feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,80 @@
|
||||
// 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
|
||||
}
|
||||
Reference in New Issue
Block a user