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