Files
tensor/signal/conv.go
T

1012 lines
34 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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
}