1012 lines
34 KiB
Go
1012 lines
34 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package signal
|
||
|
|
|
||
|
|
import (
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
|
|
)
|
||
|
|
|
||
|
|
import "sourcedock.dev/petrbalvin/tensor/internal/engine"
|
||
|
|
|
||
|
|
// 2-D convolution. The forward pass walks every output
|
||
|
|
// position and sums the contribution of every (cIn, kH, kW) input
|
||
|
|
// patch times the kernel weights. NCHW layout: input is (N, C_in,
|
||
|
|
// H, W), kernel is (C_out, C_in, kH, kW), output is (N, C_out, H_out,
|
||
|
|
// W_out). padding applies to both spatial dims (a single int for now);
|
||
|
|
// stride is also a single int. bias is added per output channel.
|
||
|
|
//
|
||
|
|
// Performance note: this is a direct O(N * C_out * C_in * H_out *
|
||
|
|
// W_out * kH * kW) implementation, not an im2col + GEMM. The kernels
|
||
|
|
// hold two constraints fixed:
|
||
|
|
//
|
||
|
|
// - Tap order. Each output element accumulates its taps in exactly
|
||
|
|
// the (cIn, kH, kW) nesting the kernels have always used; float64
|
||
|
|
// addition is not associative, so any other order changes bits.
|
||
|
|
// Loop interchange that keeps this per-element order (moving the
|
||
|
|
// output position innermost) is safe; reordering taps is not.
|
||
|
|
// - Parallel split. Work is distributed over output rows, which own
|
||
|
|
// disjoint output elements, so no split can influence a sum.
|
||
|
|
//
|
||
|
|
// Non-float inputs widen once into a scratch payload before the loops:
|
||
|
|
// FloatAt widens exactly, so the multiplied values are bit-identical
|
||
|
|
// to the per-element accessor reads.
|
||
|
|
|
||
|
|
// Conv2D performs a 2-D convolution forward pass.
|
||
|
|
func Conv2D(input, kernel, bias *core.Array, stride, padding int) (*core.Array, error) {
|
||
|
|
return conv2DImpl(input, kernel, bias, stride, padding, 1)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Conv2DGroups performs a 2-D convolution with `groups` for grouped
|
||
|
|
// (groups=2) and depthwise (groups=C_in) convolutions. When groups=1
|
||
|
|
// the call is identical to Conv2D.
|
||
|
|
func Conv2DGroups(input, kernel, bias *core.Array, stride, padding, groups int) (*core.Array, error) {
|
||
|
|
return conv2DImpl(input, kernel, bias, stride, padding, groups)
|
||
|
|
}
|
||
|
|
|
||
|
|
// conv2DImpl is the shared implementation behind Conv2D and
|
||
|
|
// Conv2DGroups. groups=1 is the standard dense convolution; groups>1
|
||
|
|
// is a depthwise convolution when groups == c_in, or a grouped
|
||
|
|
// convolution otherwise (the kernel's in-channels is c_in/groups per
|
||
|
|
// group).
|
||
|
|
func conv2DImpl(input, kernel, bias *core.Array, stride, padding, groups int) (*core.Array, error) {
|
||
|
|
if input.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("Conv2D: complex arrays are not supported")
|
||
|
|
}
|
||
|
|
if kernel.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("Conv2D: complex kernel is not supported")
|
||
|
|
}
|
||
|
|
if input.NDim() != 4 {
|
||
|
|
return nil, base.Errf("Conv2D: input must be 4-D (N, C_in, H, W), got shape %s", base.ShapeText(input.Shape()))
|
||
|
|
}
|
||
|
|
if kernel.NDim() != 4 {
|
||
|
|
return nil, base.Errf("Conv2D: kernel must be 4-D (C_out, C_in, kH, kW), got shape %s", base.ShapeText(kernel.Shape()))
|
||
|
|
}
|
||
|
|
if stride < 1 {
|
||
|
|
return nil, base.Errf("Conv2D: stride must be at least 1, got %d", stride)
|
||
|
|
}
|
||
|
|
if padding < 0 {
|
||
|
|
return nil, base.Errf("Conv2D: padding must be non-negative, got %d", padding)
|
||
|
|
}
|
||
|
|
n, cIn, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3]
|
||
|
|
cOut, cInK, kH, kW := kernel.Shape()[0], kernel.Shape()[1], kernel.Shape()[2], kernel.Shape()[3]
|
||
|
|
if groups < 1 {
|
||
|
|
return nil, base.Errf("Conv2D: groups must be at least 1, got %d", groups)
|
||
|
|
}
|
||
|
|
if cIn%groups != 0 {
|
||
|
|
return nil, base.Errf("Conv2D: input channels %d must be divisible by groups %d", cIn, groups)
|
||
|
|
}
|
||
|
|
if cOut%groups != 0 {
|
||
|
|
return nil, base.Errf("Conv2D: output channels %d must be divisible by groups %d", cOut, groups)
|
||
|
|
}
|
||
|
|
cInPerGroup := cIn / groups
|
||
|
|
cOutPerGroup := cOut / groups
|
||
|
|
if cInK != cInPerGroup {
|
||
|
|
return nil, base.Errf("Conv2D: kernel in-channels %d does not match cIn/groups %d", cInK, cInPerGroup)
|
||
|
|
}
|
||
|
|
// Guard the numerators before the division: a negative numerator
|
||
|
|
// truncates toward zero and masquerades as a 1-wide output.
|
||
|
|
hNum := hIn + 2*padding - kH
|
||
|
|
wNum := wIn + 2*padding - kW
|
||
|
|
if hNum < 0 || wNum < 0 {
|
||
|
|
return nil, base.Errf("Conv2D: kernel %dx%d with padding %d does not fit the %dx%d input", kH, kW, padding, hIn, wIn)
|
||
|
|
}
|
||
|
|
hOut := hNum/stride + 1
|
||
|
|
wOut := wNum/stride + 1
|
||
|
|
if hOut < 1 || wOut < 1 {
|
||
|
|
return nil, base.Errf("Conv2D: output is empty (H_out=%d, W_out=%d)", hOut, wOut)
|
||
|
|
}
|
||
|
|
if bias != nil && (bias.NDim() != 1 || bias.Len() != cOut) {
|
||
|
|
return nil, base.Errf("Conv2D: bias must be 1-D of length %d, got shape %s", cOut, base.ShapeText(bias.Shape()))
|
||
|
|
}
|
||
|
|
// The bias widens through FloatAt like the payloads, which reads
|
||
|
|
// only real dtypes: a complex bias is refused instead of widened.
|
||
|
|
if bias != nil && bias.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("Conv2D: complex bias is not supported")
|
||
|
|
}
|
||
|
|
out, oerr := core.Zeros(core.Float, []int{n, cOut, hOut, wOut}...)
|
||
|
|
if oerr != nil {
|
||
|
|
return nil, oerr
|
||
|
|
}
|
||
|
|
// The fast path streams raw float64 payloads; other dtypes widen
|
||
|
|
// once into scratch (exact, see widenFloats) so the loop below
|
||
|
|
// never branches on dtype. Dispatching once here keeps the
|
||
|
|
// per-element work free of accessor calls and bounds on the shape
|
||
|
|
// of the read.
|
||
|
|
inF := input.RawFloats()
|
||
|
|
if inF == nil || input.Strided() {
|
||
|
|
inF = widenFloats(input)
|
||
|
|
}
|
||
|
|
kerF := kernel.RawFloats()
|
||
|
|
if kerF == nil || kernel.Strided() {
|
||
|
|
kerF = widenFloats(kernel)
|
||
|
|
}
|
||
|
|
var biasF []float64
|
||
|
|
if bias != nil {
|
||
|
|
biasF = widenFloats(bias)
|
||
|
|
}
|
||
|
|
outF := out.RawFloats()
|
||
|
|
// Blocked over output channels: one work item is one (batch, group,
|
||
|
|
// output row, channel block) and its accumulator covers every output
|
||
|
|
// channel of the block at once. A tap's input row then feeds every
|
||
|
|
// channel of the block from the lowest cache instead of being read
|
||
|
|
// again for each of them, which is what the per-channel walk spent
|
||
|
|
// its bandwidth on. The tap walk stays (cIn, kH, kW) ascending and
|
||
|
|
// every output element is updated by exactly one tap of the walk, so
|
||
|
|
// each sum sees the same addends in the same order as before, and the
|
||
|
|
// bias is added last, as before. When the row is too wide for the
|
||
|
|
// block to stay in cache the block shrinks to one channel, which is
|
||
|
|
// the per-channel walk again.
|
||
|
|
ocBlock := max(1, min(cOutPerGroup, convAccCap/wOut))
|
||
|
|
blocksPerRow := (cOutPerGroup + ocBlock - 1) / ocBlock
|
||
|
|
// Per-channel kernel offsets within the block: the weight base of the
|
||
|
|
// block's first channel is added per item, the rest follow from the
|
||
|
|
// channel index, so a tap reads its weight without multiplying the
|
||
|
|
// channel out.
|
||
|
|
kOff := make([]int, ocBlock)
|
||
|
|
for i := range kOff {
|
||
|
|
kOff[i] = i * cInPerGroup * kH * kW
|
||
|
|
}
|
||
|
|
items := n * hOut * groups * blocksPerRow
|
||
|
|
floor := workFloorFor(wOut * ocBlock * cInPerGroup * kH * kW)
|
||
|
|
engine.ParallelMin(items, floor, func(bs, be int) {
|
||
|
|
acc := engine.GetFloat64Buf(ocBlock * wOut)
|
||
|
|
defer engine.PutFloat64Buf(acc)
|
||
|
|
for item := bs; item < be; item++ {
|
||
|
|
ob := item % blocksPerRow
|
||
|
|
rest := item / blocksPerRow
|
||
|
|
g := rest % groups
|
||
|
|
rest /= groups
|
||
|
|
oh := rest % hOut
|
||
|
|
batch := rest / hOut
|
||
|
|
oc0 := g*cOutPerGroup + ob*ocBlock
|
||
|
|
oc1 := min(oc0+ocBlock, (g+1)*cOutPerGroup)
|
||
|
|
clear(acc)
|
||
|
|
chanBase := batch*cIn*hIn*wIn + g*cInPerGroup*hIn*wIn
|
||
|
|
ocBase := oc0 * cInPerGroup * kH * kW
|
||
|
|
for icInGroup := range cInPerGroup {
|
||
|
|
inChan := chanBase + icInGroup*hIn*wIn
|
||
|
|
tapBase := icInGroup * kH * kW
|
||
|
|
for kh := range kH {
|
||
|
|
ih := oh*stride + kh - padding
|
||
|
|
if ih < 0 || ih >= hIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
inRow := inChan + ih*wIn
|
||
|
|
for kw := range kW {
|
||
|
|
// The outputs this tap reaches: ow*stride
|
||
|
|
// must land the input column inside the row.
|
||
|
|
// The reach is a property of the tap alone, so
|
||
|
|
// it is computed once for the whole block.
|
||
|
|
lo := 0
|
||
|
|
if d := padding - kw; d > 0 {
|
||
|
|
lo = (d + stride - 1) / stride
|
||
|
|
}
|
||
|
|
hi := wOut
|
||
|
|
if m := wIn - 1 + padding - kw; m < 0 {
|
||
|
|
continue
|
||
|
|
} else if t := m/stride + 1; t < hi {
|
||
|
|
hi = t
|
||
|
|
}
|
||
|
|
tap := kw - padding
|
||
|
|
if hi <= lo {
|
||
|
|
// The tap reaches no output here (its
|
||
|
|
// whole window lies in the padding): the
|
||
|
|
// unit-stride sub-slices below would get
|
||
|
|
// a high bound under their low one.
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
tapOff := tapBase + kh*kW + kw
|
||
|
|
if stride == 1 {
|
||
|
|
// Unit stride lets both sides range
|
||
|
|
// over paired sub-slices: no bounds
|
||
|
|
// checks and no index multiply in
|
||
|
|
// the hot loop.
|
||
|
|
src := inF[inRow+lo+tap : inRow+hi+tap]
|
||
|
|
for o := range oc1 - oc0 {
|
||
|
|
k := kerF[ocBase+kOff[o]+tapOff]
|
||
|
|
dst := acc[o*wOut+lo : o*wOut+hi]
|
||
|
|
for j, v := range src {
|
||
|
|
dst[j] += v * k
|
||
|
|
}
|
||
|
|
}
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
for o := range oc1 - oc0 {
|
||
|
|
k := kerF[ocBase+kOff[o]+tapOff]
|
||
|
|
base := o * wOut
|
||
|
|
for ow := lo; ow < hi; ow++ {
|
||
|
|
acc[base+ow] += inF[inRow+ow*stride+tap] * k
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
for o := range oc1 - oc0 {
|
||
|
|
oc := oc0 + o
|
||
|
|
b := 0.0
|
||
|
|
if biasF != nil {
|
||
|
|
b = biasF[oc]
|
||
|
|
}
|
||
|
|
off := (batch*cOut + oc) * hOut * wOut
|
||
|
|
base := o * wOut
|
||
|
|
for ow := range wOut {
|
||
|
|
outF[off+oh*wOut+ow] = acc[base+ow] + b
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
})
|
||
|
|
return out, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// conv1DRowBlock is the output-length grain the 1-D kernel splits rows
|
||
|
|
// into when a row alone would starve the worker pool (a single long
|
||
|
|
// signal with a handful of output channels is exactly that case).
|
||
|
|
const conv1DRowBlock = 512
|
||
|
|
|
||
|
|
// convAccCap bounds one work item's channel-block accumulator, in
|
||
|
|
// float64 entries: a block of channels·row values is what the taps of
|
||
|
|
// a whole row feed, and it stays cheap only while it fits the lowest
|
||
|
|
// cache. Above the cap the block covers fewer channels, down to one,
|
||
|
|
// where the kernel degenerates to the per-channel walk it replaced.
|
||
|
|
// Every channel-blocked kernel shares it.
|
||
|
|
const convAccCap = 1 << 12
|
||
|
|
|
||
|
|
// Conv1D performs a 1-D convolution forward pass. input is (N, C_in,
|
||
|
|
// L), kernel is (C_out, C_in, kL), output is (N, C_out, L_out). NCL
|
||
|
|
// layout.
|
||
|
|
func Conv1D(input, kernel, bias *core.Array, stride, padding, dilation int) (*core.Array, error) {
|
||
|
|
if input.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("Conv1D: complex arrays are not supported")
|
||
|
|
}
|
||
|
|
if kernel.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("Conv1D: complex kernel is not supported")
|
||
|
|
}
|
||
|
|
if input.NDim() != 3 {
|
||
|
|
return nil, base.Errf("Conv1D: input must be 3-D (N, C_in, L), got shape %s", base.ShapeText(input.Shape()))
|
||
|
|
}
|
||
|
|
if kernel.NDim() != 3 {
|
||
|
|
return nil, base.Errf("Conv1D: kernel must be 3-D (C_out, C_in, kL), got shape %s", base.ShapeText(kernel.Shape()))
|
||
|
|
}
|
||
|
|
if stride < 1 || dilation < 1 || padding < 0 {
|
||
|
|
return nil, base.Errf("Conv1D: stride/dilation must be ≥1, padding ≥0")
|
||
|
|
}
|
||
|
|
if dilation != 1 {
|
||
|
|
return nil, base.Errf("Conv1D: dilation != 1 is not implemented yet")
|
||
|
|
}
|
||
|
|
n, cIn, lIn := input.Shape()[0], input.Shape()[1], input.Shape()[2]
|
||
|
|
cOut, cInK, kL := kernel.Shape()[0], kernel.Shape()[1], kernel.Shape()[2]
|
||
|
|
if cInK != cIn {
|
||
|
|
return nil, base.Errf("Conv1D: kernel in-channels %d does not match input %d", cInK, cIn)
|
||
|
|
}
|
||
|
|
lNum := lIn + 2*padding - kL
|
||
|
|
if lNum < 0 {
|
||
|
|
return nil, base.Errf("Conv1D: kernel length %d with padding %d does not fit the length-%d input", kL, padding, lIn)
|
||
|
|
}
|
||
|
|
lOut := lNum/stride + 1
|
||
|
|
if lOut < 1 {
|
||
|
|
return nil, base.Errf("Conv1D: output is empty (L_out=%d)", lOut)
|
||
|
|
}
|
||
|
|
if bias != nil && (bias.NDim() != 1 || bias.Len() != cOut) {
|
||
|
|
return nil, base.Errf("Conv1D: bias must be 1-D of length %d, got shape %s", cOut, base.ShapeText(bias.Shape()))
|
||
|
|
}
|
||
|
|
if bias != nil && bias.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("Conv1D: complex bias is not supported")
|
||
|
|
}
|
||
|
|
out, oerr := core.Zeros(core.Float, []int{n, cOut, lOut}...)
|
||
|
|
if oerr != nil {
|
||
|
|
return nil, oerr
|
||
|
|
}
|
||
|
|
inF := input.RawFloats()
|
||
|
|
if inF == nil || input.Strided() {
|
||
|
|
inF = widenFloats(input)
|
||
|
|
}
|
||
|
|
kerF := kernel.RawFloats()
|
||
|
|
if kerF == nil || kernel.Strided() {
|
||
|
|
kerF = widenFloats(kernel)
|
||
|
|
}
|
||
|
|
var biasF []float64
|
||
|
|
if bias != nil {
|
||
|
|
biasF = widenFloats(bias)
|
||
|
|
}
|
||
|
|
outF := out.RawFloats()
|
||
|
|
// Work items are output row segments (batch, oc, a block of ol).
|
||
|
|
// Segmenting keeps a single long signal with few output channels
|
||
|
|
// from starving the pool; segments are disjoint, and each output
|
||
|
|
// accumulates its taps in the (cIn, kL) order. The per-tap spans
|
||
|
|
// are shape-only and are clipped to the row block per item.
|
||
|
|
spans := convRowSpans(kL, padding, stride, lIn, lOut)
|
||
|
|
blocks := (lOut + conv1DRowBlock - 1) / conv1DRowBlock
|
||
|
|
items := n * cOut * blocks
|
||
|
|
engine.ParallelMin(items, workFloorFor(min(conv1DRowBlock, lOut)*cIn*kL), func(bs, be int) {
|
||
|
|
acc := engine.GetFloat64Buf(min(conv1DRowBlock, lOut))
|
||
|
|
defer engine.PutFloat64Buf(acc)
|
||
|
|
for item := bs; item < be; item++ {
|
||
|
|
block := item % blocks
|
||
|
|
row := item / blocks
|
||
|
|
oc := row % cOut
|
||
|
|
batch := row / cOut
|
||
|
|
ols := block * conv1DRowBlock
|
||
|
|
ole := min(ols+conv1DRowBlock, lOut)
|
||
|
|
clear(acc)
|
||
|
|
chanBase := batch * cIn * lIn
|
||
|
|
kerBase := oc * cIn * kL
|
||
|
|
// The interior of the block reached by every tap of the
|
||
|
|
// kernel: [fa, fb) gets one fused sweep per input channel
|
||
|
|
// with the running sum held in a register, instead of one
|
||
|
|
// accumulator pass per tap. Each element still receives
|
||
|
|
// its taps in (cIn, kL) order, so the accumulated bits
|
||
|
|
// are unchanged. Outside the register-sized kernel rows
|
||
|
|
// the per-tap walk below covers the whole block.
|
||
|
|
fa := max(spans[0].lo, ols)
|
||
|
|
fb := min(spans[kL-1].hi, ole)
|
||
|
|
fused := stride == 1 && kL >= 2 && kL <= 7 && fb > fa
|
||
|
|
for ic := range cIn {
|
||
|
|
inRow := chanBase + ic*lIn
|
||
|
|
kerRow := kerBase + ic*kL
|
||
|
|
for kl := range kL {
|
||
|
|
sp := spans[kl]
|
||
|
|
if !sp.valid {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
k := kerF[kerRow+kl]
|
||
|
|
// Clip the tap's precomputed span to this row
|
||
|
|
// block and to the edges of the fused region;
|
||
|
|
// the clip only narrows.
|
||
|
|
lo := max(sp.lo, ols)
|
||
|
|
hi := min(sp.hi, ole)
|
||
|
|
if fused {
|
||
|
|
hi = min(hi, fa)
|
||
|
|
}
|
||
|
|
tap := sp.tap
|
||
|
|
if hi > lo {
|
||
|
|
if stride == 1 {
|
||
|
|
// Unit stride: paired sub-slices
|
||
|
|
// keep the hot loop free of bounds
|
||
|
|
// checks and index multiplies.
|
||
|
|
src := inF[inRow+lo+tap : inRow+hi+tap]
|
||
|
|
dst := acc[lo-ols : hi-ols]
|
||
|
|
for j, v := range src {
|
||
|
|
dst[j] += v * k
|
||
|
|
}
|
||
|
|
} else {
|
||
|
|
for ol := lo; ol < hi; ol++ {
|
||
|
|
acc[ol-ols] += inF[inRow+ol*stride+tap] * k
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
if !fused {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
lo = max(sp.lo, fb)
|
||
|
|
hi = min(sp.hi, ole)
|
||
|
|
if hi <= lo {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
src := inF[inRow+lo+tap : inRow+hi+tap]
|
||
|
|
dst := acc[lo-ols : hi-ols]
|
||
|
|
for j, v := range src {
|
||
|
|
dst[j] += v * k
|
||
|
|
}
|
||
|
|
}
|
||
|
|
if !fused {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
// The fused interior sweep: one shifted window of the
|
||
|
|
// input row, its sub-slices in lockstep, taps in
|
||
|
|
// registers.
|
||
|
|
src := inF[inRow+fa-padding : inRow+fb-padding+kL-1]
|
||
|
|
dst := acc[fa-ols : fb-ols]
|
||
|
|
m := len(dst)
|
||
|
|
switch kL {
|
||
|
|
case 2:
|
||
|
|
k0, k1 := kerF[kerRow], kerF[kerRow+1]
|
||
|
|
s1 := src[1 : 1+m]
|
||
|
|
for i, v0 := range src[:m] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
dst[i] = d + s1[i]*k1
|
||
|
|
}
|
||
|
|
case 3:
|
||
|
|
k0, k1, k2 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2]
|
||
|
|
s1 := src[1 : 1+m]
|
||
|
|
s2 := src[2 : 2+m]
|
||
|
|
for i, v0 := range src[:m] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
dst[i] = d + s2[i]*k2
|
||
|
|
}
|
||
|
|
case 4:
|
||
|
|
k0, k1, k2, k3 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3]
|
||
|
|
s1 := src[1 : 1+m]
|
||
|
|
s2 := src[2 : 2+m]
|
||
|
|
s3 := src[3 : 3+m]
|
||
|
|
for i, v0 := range src[:m] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
d += s2[i] * k2
|
||
|
|
dst[i] = d + s3[i]*k3
|
||
|
|
}
|
||
|
|
case 5:
|
||
|
|
k0, k1, k2, k3, k4 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4]
|
||
|
|
s1 := src[1 : 1+m]
|
||
|
|
s2 := src[2 : 2+m]
|
||
|
|
s3 := src[3 : 3+m]
|
||
|
|
s4 := src[4 : 4+m]
|
||
|
|
for i, v0 := range src[:m] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
d += s2[i] * k2
|
||
|
|
d += s3[i] * k3
|
||
|
|
dst[i] = d + s4[i]*k4
|
||
|
|
}
|
||
|
|
case 6:
|
||
|
|
k0, k1, k2, k3, k4, k5 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4], kerF[kerRow+5]
|
||
|
|
s1 := src[1 : 1+m]
|
||
|
|
s2 := src[2 : 2+m]
|
||
|
|
s3 := src[3 : 3+m]
|
||
|
|
s4 := src[4 : 4+m]
|
||
|
|
s5 := src[5 : 5+m]
|
||
|
|
for i, v0 := range src[:m] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
d += s2[i] * k2
|
||
|
|
d += s3[i] * k3
|
||
|
|
d += s4[i] * k4
|
||
|
|
dst[i] = d + s5[i]*k5
|
||
|
|
}
|
||
|
|
case 7:
|
||
|
|
k0, k1, k2, k3, k4, k5, k6 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4], kerF[kerRow+5], kerF[kerRow+6]
|
||
|
|
s1 := src[1 : 1+m]
|
||
|
|
s2 := src[2 : 2+m]
|
||
|
|
s3 := src[3 : 3+m]
|
||
|
|
s4 := src[4 : 4+m]
|
||
|
|
s5 := src[5 : 5+m]
|
||
|
|
s6 := src[6 : 6+m]
|
||
|
|
for i, v0 := range src[:m] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
d += s2[i] * k2
|
||
|
|
d += s3[i] * k3
|
||
|
|
d += s4[i] * k4
|
||
|
|
d += s5[i] * k5
|
||
|
|
dst[i] = d + s6[i]*k6
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
b := 0.0
|
||
|
|
if biasF != nil {
|
||
|
|
b = biasF[oc]
|
||
|
|
}
|
||
|
|
off := batch*cOut*lOut + oc*lOut
|
||
|
|
for ol := ols; ol < ole; ol++ {
|
||
|
|
outF[off+ol] = acc[ol-ols] + b
|
||
|
|
}
|
||
|
|
}
|
||
|
|
})
|
||
|
|
return out, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// convGatherSpan is the reach of one transposed-convolution kernel
|
||
|
|
// column: the outputs it feeds sit stride apart on the phase lattice,
|
||
|
|
// at lattice indices [j0, j0+count), and read the input row from
|
||
|
|
// base on. first is the lowest output column reached and phase is its
|
||
|
|
// residue modulo the stride, which decides which accumulator partition
|
||
|
|
// the tap lands in. A tap with an empty reach is marked invalid.
|
||
|
|
type convGatherSpan struct {
|
||
|
|
first, j0, count, base, phase int
|
||
|
|
valid bool
|
||
|
|
}
|
||
|
|
|
||
|
|
// convGatherSpans precomputes the per-tap gather spans for one
|
||
|
|
// transposed kernel row. The values depend on the kernel column and
|
||
|
|
// the shape alone, so the hot loop reads them from this table instead
|
||
|
|
// of dividing per output position.
|
||
|
|
func convGatherSpans(kW, padding, stride, wIn, wOut int) []convGatherSpan {
|
||
|
|
spans := make([]convGatherSpan, kW)
|
||
|
|
for kw := range kW {
|
||
|
|
lo := 0
|
||
|
|
if d := kw - padding; d > 0 {
|
||
|
|
lo = d
|
||
|
|
}
|
||
|
|
hi := wOut
|
||
|
|
if m := wIn*stride + kw - padding; m < hi {
|
||
|
|
hi = m
|
||
|
|
}
|
||
|
|
if hi <= lo {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
r := (kw - padding - lo) % stride
|
||
|
|
if r < 0 {
|
||
|
|
r += stride
|
||
|
|
}
|
||
|
|
first := lo + r
|
||
|
|
count := (hi - first + stride - 1) / stride
|
||
|
|
if count <= 0 {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
spans[kw].first = first
|
||
|
|
spans[kw].count = count
|
||
|
|
spans[kw].j0 = first / stride
|
||
|
|
spans[kw].base = (first + padding - kw) / stride
|
||
|
|
spans[kw].phase = first % stride
|
||
|
|
spans[kw].valid = true
|
||
|
|
}
|
||
|
|
return spans
|
||
|
|
}
|
||
|
|
|
||
|
|
// convTapSpan is the reach of one kernel tap along an output row: the
|
||
|
|
// tap lands on outputs [lo, hi) and reads the input row offset by tap
|
||
|
|
// from them. A tap whose whole window falls into the padding is marked
|
||
|
|
// invalid, so the hot loop skips it without recomputing why.
|
||
|
|
type convTapSpan struct {
|
||
|
|
lo, hi, tap int
|
||
|
|
valid bool
|
||
|
|
}
|
||
|
|
|
||
|
|
// convRowSpans precomputes the per-tap spans of one kernel row for
|
||
|
|
// stride and padding. The bounds depend on the kernel column and the
|
||
|
|
// shape alone, never on the output position, so the innermost loops
|
||
|
|
// load them from this table instead of dividing once per tap.
|
||
|
|
func convRowSpans(kW, padW, stride, wIn, wOut int) []convTapSpan {
|
||
|
|
spans := make([]convTapSpan, kW)
|
||
|
|
for kw := range kW {
|
||
|
|
lo := 0
|
||
|
|
if d := padW - kw; d > 0 {
|
||
|
|
lo = (d + stride - 1) / stride
|
||
|
|
}
|
||
|
|
// lo and tap are stored whatever the reach: the fused region
|
||
|
|
// bounds read spans[0].lo and spans[kW-1].hi, so an extreme tap
|
||
|
|
// that reaches nothing must leave an empty window (hi 0) rather
|
||
|
|
// than a zero lo beside a stale hi.
|
||
|
|
spans[kw].lo = lo
|
||
|
|
spans[kw].tap = kw - padW
|
||
|
|
if m := wIn - 1 + padW - kw; m < 0 {
|
||
|
|
continue
|
||
|
|
} else if t := m/stride + 1; t < wOut {
|
||
|
|
spans[kw].hi = t
|
||
|
|
} else {
|
||
|
|
spans[kw].hi = wOut
|
||
|
|
}
|
||
|
|
if spans[kw].hi <= lo {
|
||
|
|
// The tap reaches no output here (its whole window
|
||
|
|
// lies in the padding): the unit-stride sub-slices
|
||
|
|
// would get a high bound under their low one.
|
||
|
|
spans[kw].hi = 0
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
spans[kw].valid = true
|
||
|
|
}
|
||
|
|
return spans
|
||
|
|
}
|
||
|
|
|
||
|
|
// Conv3D performs a 3-D convolution forward pass. input is (N, C_in,
|
||
|
|
// D, H, W), kernel is (C_out, C_in, kD, kH, kW), output is (N, C_out,
|
||
|
|
// D_out, H_out, W_out). NCDHW layout.
|
||
|
|
func Conv3D(input, kernel, bias *core.Array, stride int, padding, dilation [3]int) (*core.Array, error) {
|
||
|
|
if input.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("Conv3D: complex arrays are not supported")
|
||
|
|
}
|
||
|
|
if kernel.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("Conv3D: complex kernel is not supported")
|
||
|
|
}
|
||
|
|
if input.NDim() != 5 {
|
||
|
|
return nil, base.Errf("Conv3D: input must be 5-D, got shape %s", base.ShapeText(input.Shape()))
|
||
|
|
}
|
||
|
|
if kernel.NDim() != 5 {
|
||
|
|
return nil, base.Errf("Conv3D: kernel must be 5-D (C_out, C_in, kD, kH, kW), got shape %s", base.ShapeText(kernel.Shape()))
|
||
|
|
}
|
||
|
|
if stride < 1 {
|
||
|
|
return nil, base.Errf("Conv3D: stride must be at least 1, got %d", stride)
|
||
|
|
}
|
||
|
|
for d, dil := range dilation {
|
||
|
|
if dil != 1 {
|
||
|
|
return nil, base.Errf("Conv3D: dilation != 1 not implemented yet (dim %d)", d)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
n, cIn, dIn, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3], input.Shape()[4]
|
||
|
|
cOut, cInK, kD, kH, kW := kernel.Shape()[0], kernel.Shape()[1], kernel.Shape()[2], kernel.Shape()[3], kernel.Shape()[4]
|
||
|
|
if cInK != cIn {
|
||
|
|
return nil, base.Errf("Conv3D: kernel in-channels %d does not match input %d", cInK, cIn)
|
||
|
|
}
|
||
|
|
padD, padH, padW := padding[0], padding[1], padding[2]
|
||
|
|
// Guard the numerators before the division: a negative numerator
|
||
|
|
// truncates toward zero and masquerades as a 1-wide output.
|
||
|
|
dNum := dIn + 2*padD - kD
|
||
|
|
hNum := hIn + 2*padH - kH
|
||
|
|
wNum := wIn + 2*padW - kW
|
||
|
|
if dNum < 0 || hNum < 0 || wNum < 0 {
|
||
|
|
return nil, base.Errf("Conv3D: kernel %dx%dx%d with padding %v does not fit the input", kD, kH, kW, padding)
|
||
|
|
}
|
||
|
|
dOut := dNum/stride + 1
|
||
|
|
hOut := hNum/stride + 1
|
||
|
|
wOut := wNum/stride + 1
|
||
|
|
if dOut < 1 || hOut < 1 || wOut < 1 {
|
||
|
|
return nil, base.Errf("Conv3D: output is empty")
|
||
|
|
}
|
||
|
|
if bias != nil && (bias.NDim() != 1 || bias.Len() != cOut) {
|
||
|
|
return nil, base.Errf("Conv3D: bias must be 1-D of length %d, got shape %s", cOut, base.ShapeText(bias.Shape()))
|
||
|
|
}
|
||
|
|
if bias != nil && bias.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("Conv3D: complex bias is not supported")
|
||
|
|
}
|
||
|
|
out, oerr := core.Zeros(core.Float, []int{n, cOut, dOut, hOut, wOut}...)
|
||
|
|
if oerr != nil {
|
||
|
|
return nil, oerr
|
||
|
|
}
|
||
|
|
strideHW := wIn
|
||
|
|
strideDHW := hIn * wIn
|
||
|
|
strideCDHW := dIn * hIn * wIn
|
||
|
|
inF := input.RawFloats()
|
||
|
|
if inF == nil || input.Strided() {
|
||
|
|
inF = widenFloats(input)
|
||
|
|
}
|
||
|
|
kerF := kernel.RawFloats()
|
||
|
|
if kerF == nil || kernel.Strided() {
|
||
|
|
kerF = widenFloats(kernel)
|
||
|
|
}
|
||
|
|
var biasF []float64
|
||
|
|
if bias != nil {
|
||
|
|
biasF = widenFloats(bias)
|
||
|
|
}
|
||
|
|
outF := out.RawFloats()
|
||
|
|
// One work item per output depth-height row (batch, oc, od, oh).
|
||
|
|
// Rows are disjoint; the taps of one output accumulate in the
|
||
|
|
// (cIn, kD, kH, kW) order. The channel-blocked shape was tried
|
||
|
|
// here and measured: it paid 12 to 14 percent at 32 and 128
|
||
|
|
// channels and lost 12 percent at 4, where the per-tap channel
|
||
|
|
// loop costs more than the re-reads it saves, so the per-channel
|
||
|
|
// walk stays. The per-tap column spans are shape-only, so they
|
||
|
|
// are read from a table built once per call.
|
||
|
|
spans := convRowSpans(kW, padW, stride, wIn, wOut)
|
||
|
|
items := n * cOut * dOut * hOut
|
||
|
|
engine.ParallelMin(items, workFloorFor(wOut*cIn*kD*kH*kW), func(bs, be int) {
|
||
|
|
acc := engine.GetFloat64Buf(wOut)
|
||
|
|
defer engine.PutFloat64Buf(acc)
|
||
|
|
for item := bs; item < be; item++ {
|
||
|
|
oh := item % hOut
|
||
|
|
t := item / hOut
|
||
|
|
od := t % dOut
|
||
|
|
t /= dOut
|
||
|
|
oc := t % cOut
|
||
|
|
batch := t / cOut
|
||
|
|
clear(acc)
|
||
|
|
for ic := range cIn {
|
||
|
|
inChan := batch*cIn*strideCDHW + ic*strideCDHW
|
||
|
|
kerChan := oc*cInK*kD*kH*kW + ic*kD*kH*kW
|
||
|
|
for kd := range kD {
|
||
|
|
id := od*stride + kd - padD
|
||
|
|
if id < 0 || id >= dIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
inDepth := inChan + id*strideDHW
|
||
|
|
kerDepth := kerChan + kd*kH*kW
|
||
|
|
for kh := range kH {
|
||
|
|
ih := oh*stride + kh - padH
|
||
|
|
if ih < 0 || ih >= hIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
inRow := inDepth + ih*strideHW
|
||
|
|
kerRow := kerDepth + kh*kW
|
||
|
|
if stride == 1 && kW >= 2 && kW <= 7 {
|
||
|
|
// Unit stride with a
|
||
|
|
// register-sized kernel row. The
|
||
|
|
// interior [a, b) is reached by
|
||
|
|
// every tap, so one sweep adds all
|
||
|
|
// taps there with the running sum
|
||
|
|
// held in a register instead of
|
||
|
|
// re-reading the accumulator row
|
||
|
|
// once per tap. Each element still
|
||
|
|
// receives its taps in kw order, so
|
||
|
|
// the accumulated bits are unchanged.
|
||
|
|
a := spans[0].lo
|
||
|
|
b := spans[kW-1].hi
|
||
|
|
if b <= a {
|
||
|
|
// No interior: an empty window
|
||
|
|
// hands every tap's span to the
|
||
|
|
// second part below in one
|
||
|
|
// piece, exactly once.
|
||
|
|
a, b = 0, 0
|
||
|
|
}
|
||
|
|
for kw := range kW {
|
||
|
|
sp := spans[kw]
|
||
|
|
if !sp.valid {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
k := kerF[kerRow+kw]
|
||
|
|
lo := sp.lo
|
||
|
|
hi := min(sp.hi, a)
|
||
|
|
if hi > lo {
|
||
|
|
src := inF[inRow+lo+sp.tap : inRow+hi+sp.tap]
|
||
|
|
dst := acc[lo:hi]
|
||
|
|
for j, v := range src {
|
||
|
|
dst[j] += v * k
|
||
|
|
}
|
||
|
|
}
|
||
|
|
lo = max(sp.lo, b)
|
||
|
|
hi = sp.hi
|
||
|
|
if hi > lo {
|
||
|
|
src := inF[inRow+lo+sp.tap : inRow+hi+sp.tap]
|
||
|
|
dst := acc[lo:hi]
|
||
|
|
for j, v := range src {
|
||
|
|
dst[j] += v * k
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
if b <= a {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
// The interior sweep. The taps
|
||
|
|
// live in registers; the sub-slices
|
||
|
|
// of one shifted window keep the
|
||
|
|
// loop free of bounds checks and
|
||
|
|
// index multiplies.
|
||
|
|
src := inF[inRow+a-padW : inRow+b-padW+kW-1]
|
||
|
|
dst := acc[a:b]
|
||
|
|
n := len(dst)
|
||
|
|
switch kW {
|
||
|
|
case 2:
|
||
|
|
k0, k1 := kerF[kerRow], kerF[kerRow+1]
|
||
|
|
s1 := src[1 : 1+n]
|
||
|
|
for i, v0 := range src[:n] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
dst[i] = d + s1[i]*k1
|
||
|
|
}
|
||
|
|
case 3:
|
||
|
|
k0, k1, k2 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2]
|
||
|
|
s1 := src[1 : 1+n]
|
||
|
|
s2 := src[2 : 2+n]
|
||
|
|
for i, v0 := range src[:n] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
dst[i] = d + s2[i]*k2
|
||
|
|
}
|
||
|
|
case 4:
|
||
|
|
k0, k1, k2, k3 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3]
|
||
|
|
s1 := src[1 : 1+n]
|
||
|
|
s2 := src[2 : 2+n]
|
||
|
|
s3 := src[3 : 3+n]
|
||
|
|
for i, v0 := range src[:n] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
d += s2[i] * k2
|
||
|
|
dst[i] = d + s3[i]*k3
|
||
|
|
}
|
||
|
|
case 5:
|
||
|
|
k0, k1, k2, k3, k4 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4]
|
||
|
|
s1 := src[1 : 1+n]
|
||
|
|
s2 := src[2 : 2+n]
|
||
|
|
s3 := src[3 : 3+n]
|
||
|
|
s4 := src[4 : 4+n]
|
||
|
|
for i, v0 := range src[:n] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
d += s2[i] * k2
|
||
|
|
d += s3[i] * k3
|
||
|
|
dst[i] = d + s4[i]*k4
|
||
|
|
}
|
||
|
|
case 6:
|
||
|
|
k0, k1, k2, k3, k4, k5 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4], kerF[kerRow+5]
|
||
|
|
s1 := src[1 : 1+n]
|
||
|
|
s2 := src[2 : 2+n]
|
||
|
|
s3 := src[3 : 3+n]
|
||
|
|
s4 := src[4 : 4+n]
|
||
|
|
s5 := src[5 : 5+n]
|
||
|
|
for i, v0 := range src[:n] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
d += s2[i] * k2
|
||
|
|
d += s3[i] * k3
|
||
|
|
d += s4[i] * k4
|
||
|
|
dst[i] = d + s5[i]*k5
|
||
|
|
}
|
||
|
|
case 7:
|
||
|
|
k0, k1, k2, k3, k4, k5, k6 := kerF[kerRow], kerF[kerRow+1], kerF[kerRow+2], kerF[kerRow+3], kerF[kerRow+4], kerF[kerRow+5], kerF[kerRow+6]
|
||
|
|
s1 := src[1 : 1+n]
|
||
|
|
s2 := src[2 : 2+n]
|
||
|
|
s3 := src[3 : 3+n]
|
||
|
|
s4 := src[4 : 4+n]
|
||
|
|
s5 := src[5 : 5+n]
|
||
|
|
s6 := src[6 : 6+n]
|
||
|
|
for i, v0 := range src[:n] {
|
||
|
|
d := dst[i] + v0*k0
|
||
|
|
d += s1[i] * k1
|
||
|
|
d += s2[i] * k2
|
||
|
|
d += s3[i] * k3
|
||
|
|
d += s4[i] * k4
|
||
|
|
d += s5[i] * k5
|
||
|
|
dst[i] = d + s6[i]*k6
|
||
|
|
}
|
||
|
|
}
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
for kw := range kW {
|
||
|
|
sp := spans[kw]
|
||
|
|
if !sp.valid {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
k := kerF[kerRow+kw]
|
||
|
|
for ow := sp.lo; ow < sp.hi; ow++ {
|
||
|
|
acc[ow] += inF[inRow+ow*stride+sp.tap] * k
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
b := 0.0
|
||
|
|
if biasF != nil {
|
||
|
|
b = biasF[oc]
|
||
|
|
}
|
||
|
|
off := ((batch*cOut+oc)*dOut+od)*hOut*wOut + oh*wOut
|
||
|
|
for ow := range wOut {
|
||
|
|
outF[off+ow] = acc[ow] + b
|
||
|
|
}
|
||
|
|
}
|
||
|
|
})
|
||
|
|
return out, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// ConvTranspose2D performs the transposed convolution (sometimes
|
||
|
|
// called "deconvolution"). Input (N, C_in, H, W), kernel (C_in,
|
||
|
|
// C_out, kH, kW), output (N, C_out, H_out, W_out) where H_out =
|
||
|
|
// (H-1)*stride - 2*padding + kH (and analogously for W). bias is
|
||
|
|
// optional, applied per output channel.
|
||
|
|
func ConvTranspose2D(input, kernel, bias *core.Array, stride, padding int) (*core.Array, error) {
|
||
|
|
if input.Dtype() == core.Complex || kernel.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("ConvTranspose2D: complex arrays are not supported")
|
||
|
|
}
|
||
|
|
if input.NDim() != 4 {
|
||
|
|
return nil, base.Errf("ConvTranspose2D: input must be 4-D (N, C_in, H, W), got shape %s", base.ShapeText(input.Shape()))
|
||
|
|
}
|
||
|
|
if kernel.NDim() != 4 {
|
||
|
|
return nil, base.Errf("ConvTranspose2D: kernel must be 4-D (C_in, C_out, kH, kW), got shape %s", base.ShapeText(kernel.Shape()))
|
||
|
|
}
|
||
|
|
if stride < 1 {
|
||
|
|
return nil, base.Errf("ConvTranspose2D: stride must be at least 1, got %d", stride)
|
||
|
|
}
|
||
|
|
if padding < 0 {
|
||
|
|
return nil, base.Errf("ConvTranspose2D: padding must be non-negative, got %d", padding)
|
||
|
|
}
|
||
|
|
n, cIn, hIn, wIn := input.Shape()[0], input.Shape()[1], input.Shape()[2], input.Shape()[3]
|
||
|
|
cInK, cOut, kH, kW := kernel.Shape()[0], kernel.Shape()[1], kernel.Shape()[2], kernel.Shape()[3]
|
||
|
|
if cInK != cIn {
|
||
|
|
return nil, base.Errf("ConvTranspose2D: kernel in-channels %d does not match input %d", cInK, cIn)
|
||
|
|
}
|
||
|
|
hOut := (hIn-1)*stride - 2*padding + kH
|
||
|
|
wOut := (wIn-1)*stride - 2*padding + kW
|
||
|
|
if hOut < 1 || wOut < 1 {
|
||
|
|
return nil, base.Errf("ConvTranspose2D: output is empty")
|
||
|
|
}
|
||
|
|
if bias != nil && (bias.NDim() != 1 || bias.Len() != cOut) {
|
||
|
|
return nil, base.Errf("ConvTranspose2D: bias must be 1-D of length %d, got shape %s", cOut, base.ShapeText(bias.Shape()))
|
||
|
|
}
|
||
|
|
if bias != nil && bias.Dtype() == core.Complex {
|
||
|
|
return nil, base.Errf("ConvTranspose2D: complex bias is not supported")
|
||
|
|
}
|
||
|
|
out, oerr := core.Zeros(core.Float, []int{n, cOut, hOut, wOut}...)
|
||
|
|
if oerr != nil {
|
||
|
|
return nil, oerr
|
||
|
|
}
|
||
|
|
inF := input.RawFloats()
|
||
|
|
if inF == nil || input.Strided() {
|
||
|
|
inF = widenFloats(input)
|
||
|
|
}
|
||
|
|
kerF := kernel.RawFloats()
|
||
|
|
if kerF == nil || kernel.Strided() {
|
||
|
|
kerF = widenFloats(kernel)
|
||
|
|
}
|
||
|
|
var biasF []float64
|
||
|
|
if bias != nil {
|
||
|
|
biasF = widenFloats(bias)
|
||
|
|
}
|
||
|
|
outF := out.RawFloats()
|
||
|
|
// The scatter form (walking input positions and adding into every
|
||
|
|
// output they touch) becomes a gather here: each output row sums
|
||
|
|
// the input positions that map onto it. Blocked over output
|
||
|
|
// channels, as Conv2D is: one work item is one (batch, output row,
|
||
|
|
// channel block) and its accumulator holds every output channel of
|
||
|
|
// the block at once, so a gathered input value is spent on every
|
||
|
|
// channel of the block instead of being gathered again for each of
|
||
|
|
// them. For a single output the contributions still arrive in the
|
||
|
|
// (cIn, kH, kW) order the scatter visited them, and the bias is
|
||
|
|
// added last, so the accumulated bits are unchanged. Rows are
|
||
|
|
// disjoint, so the parallel split cannot touch a sum.
|
||
|
|
//
|
||
|
|
// The accumulator is partitioned by phase: a tap with reversed
|
||
|
|
// column kw - padding only ever touches outputs congruent to it
|
||
|
|
// modulo the stride, so each phase owns a contiguous sub-lattice
|
||
|
|
// of every output row. A tap's gather then reads a contiguous
|
||
|
|
// input window into a contiguous accumulator window, with no
|
||
|
|
// divisions and no strided stores in the hot loop.
|
||
|
|
ocBlock := max(1, min(cOut, convAccCap/wOut))
|
||
|
|
blocksPerRow := (cOut + ocBlock - 1) / ocBlock
|
||
|
|
rowCap := (wOut + stride - 1) / stride
|
||
|
|
spans := convGatherSpans(kW, padding, stride, wIn, wOut)
|
||
|
|
items := n * hOut * blocksPerRow
|
||
|
|
engine.ParallelMin(items, workFloorFor(wOut*ocBlock*cIn*kH*kW), func(bs, be int) {
|
||
|
|
acc := engine.GetFloat64Buf(ocBlock * stride * rowCap)
|
||
|
|
defer engine.PutFloat64Buf(acc)
|
||
|
|
var ihMap [16]int
|
||
|
|
mapped := min(kH, len(ihMap))
|
||
|
|
for item := bs; item < be; item++ {
|
||
|
|
ob := item % blocksPerRow
|
||
|
|
rest := item / blocksPerRow
|
||
|
|
oh := rest % hOut
|
||
|
|
batch := rest / hOut
|
||
|
|
oc0 := ob * ocBlock
|
||
|
|
oc1 := min(oc0+ocBlock, cOut)
|
||
|
|
clear(acc)
|
||
|
|
// Map (oh) back to the input rows once per item: the
|
||
|
|
// scatter only paired positions with oh + padding - kh
|
||
|
|
// divisible by the stride and inside the input, and the
|
||
|
|
// mapping does not depend on the channel.
|
||
|
|
for kh := range mapped {
|
||
|
|
ih := -1
|
||
|
|
if ihNum := oh + padding - kh; ihNum%stride == 0 {
|
||
|
|
if t := ihNum / stride; t >= 0 && t < hIn {
|
||
|
|
ih = t
|
||
|
|
}
|
||
|
|
}
|
||
|
|
ihMap[kh] = ih
|
||
|
|
}
|
||
|
|
for ic := range cIn {
|
||
|
|
inChan := batch*cIn*hIn*wIn + ic*hIn*wIn
|
||
|
|
icBase := ic * cOut * kH * kW
|
||
|
|
for kh := range kH {
|
||
|
|
var ih int
|
||
|
|
if kh < mapped {
|
||
|
|
ih = ihMap[kh]
|
||
|
|
} else {
|
||
|
|
ihNum := oh + padding - kh
|
||
|
|
if ihNum%stride != 0 {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
ih = ihNum / stride
|
||
|
|
if ih < 0 || ih >= hIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
}
|
||
|
|
if ih < 0 {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
inRow := inChan + ih*wIn
|
||
|
|
khOff := icBase + kh*kW
|
||
|
|
for kw := range kW {
|
||
|
|
sp := spans[kw]
|
||
|
|
if !sp.valid {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
tapOff := khOff + oc0*kH*kW + kw
|
||
|
|
src := inF[inRow+sp.base : inRow+sp.base+sp.count]
|
||
|
|
for o := range oc1 - oc0 {
|
||
|
|
k := kerF[tapOff+o*kH*kW]
|
||
|
|
off := (o*stride + sp.phase) * rowCap
|
||
|
|
dst := acc[off+sp.j0 : off+sp.j0+sp.count]
|
||
|
|
for j, v := range src {
|
||
|
|
dst[j] += v * k
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
for o := range oc1 - oc0 {
|
||
|
|
oc := oc0 + o
|
||
|
|
b := 0.0
|
||
|
|
if biasF != nil {
|
||
|
|
b = biasF[oc]
|
||
|
|
}
|
||
|
|
off := (batch*cOut + oc) * hOut * wOut
|
||
|
|
for p := range stride {
|
||
|
|
phaseBase := (o*stride + p) * rowCap
|
||
|
|
j := 0
|
||
|
|
for ow := p; ow < wOut; ow += stride {
|
||
|
|
outF[off+oh*wOut+ow] = acc[phaseBase+j] + b
|
||
|
|
j++
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
})
|
||
|
|
return out, nil
|
||
|
|
}
|