81 lines
2.3 KiB
Go
81 lines
2.3 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package grad
|
|
|
|
import "sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
|
|
// Slice extracts a range along the given dimension as a new tensor;
|
|
// the backward writes the incoming gradient into the corresponding
|
|
// region of the original shape.
|
|
func (t *Tensor) Slice(dim, start, stop int) (*Tensor, error) {
|
|
if err := t.checkDiff("Slice"); err != nil {
|
|
return nil, err
|
|
}
|
|
out, err := core.Slice(t.data, dim, start, stop)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
orig := t.data.Shape()
|
|
dt := t.data.Dtype()
|
|
return t.unaryResult("Slice", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
|
// The narrowing Concat applies: a complex gradient reaching a
|
|
// real slice narrows by 2·Re before the span is copied. A
|
|
// slice's output dtype equals its input's, so only a complex
|
|
// gradient on a real tensor can differ here.
|
|
gn, err := narrowGradient(g, dt)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
da := gradSlot{arr: ar.borrowGrad(dt, orig), sh: orig}
|
|
outer := 1
|
|
for d := range dim {
|
|
outer *= orig[d]
|
|
}
|
|
inner := 1
|
|
for d := dim + 1; d < len(orig); d++ {
|
|
inner *= orig[d]
|
|
}
|
|
nS := stop - start
|
|
// Each kept row is one contiguous inner run, so matching dtypes
|
|
// ride raw slice moves instead of per-element accessor calls.
|
|
fast := !gn.arr.Strided() && gn.arr.Dtype() == dt && dt != core.Int
|
|
for o := range outer {
|
|
for si := range nS {
|
|
d := o*orig[dim]*inner + (start+si)*inner
|
|
s := o*nS*inner + si*inner
|
|
if fast {
|
|
copySegRaw(da.arr, gn.arr, d, s, inner)
|
|
continue
|
|
}
|
|
for j := range inner {
|
|
copyElem(da.arr, d+j, gn.arr, s+j)
|
|
}
|
|
}
|
|
}
|
|
dst[0] = da
|
|
return nil
|
|
}), nil
|
|
}
|
|
|
|
// Reshape returns a new view-equivalent tensor of the given shape; the
|
|
// backward simply reshapes the incoming gradient back.
|
|
func (t *Tensor) Reshape(shape ...int) (*Tensor, error) {
|
|
if err := t.checkDiff("Reshape"); err != nil {
|
|
return nil, err
|
|
}
|
|
out, err := core.Reshape(t.data, shape...)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
orig := append([]int{}, t.data.Shape()...)
|
|
return t.unaryResult("Reshape", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
|
gr, err := core.Reshape(g.arr, orig...)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
dst[0] = gradSlot{arr: gr, sh: orig}
|
|
return nil
|
|
}), nil
|
|
}
|