feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,42 @@
|
||||
// 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"
|
||||
)
|
||||
|
||||
// TransposeAxes reorders the axes of a tensor by dims. The backward
|
||||
// applies the inverse permutation to the incoming gradient: axis moves
|
||||
// are invertible data motion, so no element mixing occurs and the
|
||||
// gradient is exactly the same move played backwards.
|
||||
func (t *Tensor) TransposeAxes(dims ...int) (*Tensor, error) {
|
||||
if err := t.checkDiff("TransposeAxes"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out, err := core.TransposeAxes(t.data, dims...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
perm := append([]int(nil), dims...)
|
||||
orig := t.data.Shape()
|
||||
return t.unaryResult("TransposeAxes", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
||||
dx, err := core.TransposeAxes(g.arr, inversePerm(perm)...)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dst[0] = gradSlot{arr: dx, sh: orig}
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
|
||||
// inversePerm flips an axis permutation: if out = permute(x, p), then
|
||||
// permute(out, p⁻¹) restores x's axis order.
|
||||
func inversePerm(perm []int) []int {
|
||||
inv := make([]int, len(perm))
|
||||
for i, p := range perm {
|
||||
inv[p] = i
|
||||
}
|
||||
return inv
|
||||
}
|
||||
Reference in New Issue
Block a user