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