43 lines
1.2 KiB
Go
43 lines
1.2 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"
|
||
|
|
)
|
||
|
|
|
||
|
|
// 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
|
||
|
|
}
|