// Copyright (c) 2026 Petr Balvín (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 }