Files
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

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
}