tests/HumanoidRetargeter.Tests/Dl/SameGoldenTests.cs

Unit tests for the SAME model port. Verifies parsing of weights blobs, encoder/decoder forward outputs, feature pipeline and whole pipeline against golden tensors produced by a verified Python pipeline, using numeric tolerances.

File Access
using HumanoidRetargeter.Core.Dl;
using Xunit;

namespace HumanoidRetargeter.Tests.Dl;

/// <summary>
/// Parity of the pure-managed SAME port against golden tensors produced by the verified
/// Python pipeline (dev/m10/scripts/export_golden.py: SAME ckpt0 + the explicit-math GAT
/// wrappers proven bit-equal to the original PyG layers). Gate: ≤ 1e-3 everywhere.
/// </summary>
public class SameGoldenTests
{
    private const float ModelTolerance = 1e-3f;

    // ---------------------------------------------------------------- weights blob

    [Fact]
    public void WeightsBlobParsesWithExpectedArchitecture()
    {
        var weights = DlFixtures.Weights.Value;
        Assert.Equal(48, weights.Tensors.Count); // 8 GAT layers × 4 tensors + 16 ms stats

        var first = weights.Get("encoder.0.convs.0.lin_src.weight");
        Assert.Equal(new[] { 256, 32 }, first.Shape);
        var lastDec = weights.Get("decoder.0.deconvs.3.lin_src.weight");
        Assert.Equal(new[] { 11, 262 }, lastDec.Shape);
        Assert.Equal(6, weights.Stat("q_m").Length);
        Assert.Equal(4, weights.Stat("r_s").Length);
    }

    // ---------------------------------------------------------------- model forward parity

    [Fact]
    public void EncoderMatchesGoldenOnGoldenInputs()
    {
        var golden = DlFixtures.Golden.Value;
        var model = DlFixtures.Model.Value;

        var x = golden.Get("src.x").Data;
        var edges = golden.Get("src.edge_index");
        var e = edges.Shape[1];
        var edgeSrc = DlFixtures.ToInts(edges.Data.Take(e).ToArray());
        var edgeDst = DlFixtures.ToInts(edges.Data.Skip(e).ToArray());
        var batch = DlFixtures.ToInts(golden.Get("src.batch").Data);
        var frames = batch.Max() + 1;

        var nodeZ = model.EncodeNodes(x, edgeSrc, edgeDst);
        AssertClose(golden.Get("enc.node_z").Data, nodeZ, ModelTolerance, "node_z");

        var z = SameModel.MaxPool(nodeZ, SameModel.ZDim, batch, frames);
        AssertClose(golden.Get("enc.z").Data, z, ModelTolerance, "z");
    }

    [Fact]
    public void DecoderMatchesGoldenOnGoldenInputs()
    {
        var golden = DlFixtures.Golden.Value;
        var model = DlFixtures.Model.Value;
        var (_, _, frames) = DlFixtures.GoldenNames();

        var z = golden.Get("enc.z").Data;
        var tgtX = golden.Get("tgt.x");
        var j = tgtX.Shape[0];
        var edges = golden.Get("tgt.edge_index");
        var e = edges.Shape[1];
        var src1 = DlFixtures.ToInts(edges.Data.Take(e).ToArray());
        var dst1 = DlFixtures.ToInts(edges.Data.Skip(e).ToArray());

        // Tile the single-frame golden graph across the batch.
        var x = new float[frames * tgtX.Data.Length];
        var batch = new int[frames * j];
        var edgeSrc = new int[frames * e];
        var edgeDst = new int[frames * e];
        for (var f = 0; f < frames; f++)
        {
            Array.Copy(tgtX.Data, 0, x, f * tgtX.Data.Length, tgtX.Data.Length);
            for (var i = 0; i < j; i++)
                batch[f * j + i] = f;
            for (var i = 0; i < e; i++)
            {
                edgeSrc[f * e + i] = src1[i] + f * j;
                edgeDst[f * e + i] = dst1[i] + f * j;
            }
        }

        var hatD = model.Decode(z, x, edgeSrc, edgeDst, batch);
        AssertClose(golden.Get("dec.hatD").Data, hatD, ModelTolerance, "hatD");
    }

    // ---------------------------------------------------------------- feature pipeline parity

    [Fact]
    public void SourceFeaturesMatchGolden()
    {
        var golden = DlFixtures.Golden.Value;
        var (srcNames, _, frames) = DlFixtures.GoldenNames();

        var scene = DlFixtures.FixtureBvhTruncated(frames + 2); // tpose + wov frames
        var graph = SameFeatures.BuildSourceGraph(
            scene, 0, map: null, DlFixtures.Stats.Value,
            new SameFeatures.SourceOptions
            {
                NativeFrameDrop = true,
                Align = false,
                GroundShift = false,
            });

        Assert.Equal(frames, graph.FrameCount);
        Assert.Equal(srcNames.Length, graph.JointCount);

        var map = DlFixtures.MapNames(srcNames, graph.JointNames);
        var goldenX = golden.Get("src.x").Data;
        var width = SameModel.InputDim;

        // The importer resamples the 0.0333333s-frame-time file onto an exact 30 fps grid
        // (≈1e-5 of slerp per frame) and works in float32 against numpy's float64;
        // measured worst deviation 2.6e-5.
        const float tolerance = 1e-3f;
        var worst = 0f;
        var worstWhere = "";
        for (var f = 0; f < frames; f++)
        {
            for (var g = 0; g < srcNames.Length; g++)
            {
                var ours = (f * graph.JointCount + map[g]) * width;
                var theirs = (f * srcNames.Length + g) * width;
                for (var c = 0; c < width; c++)
                {
                    var d = MathF.Abs(goldenX[theirs + c] - graph.X[ours + c]);
                    if (d > worst)
                    {
                        worst = d;
                        worstWhere = $"frame {f}, joint {srcNames[g]}, column {c}";
                    }
                }
            }
        }
        Assert.True(worst <= tolerance,
            $"Source features deviate from golden by {worst} (> {tolerance}) at {worstWhere}.");
    }

    [Fact]
    public void TargetGraphMatchesGolden()
    {
        var golden = DlFixtures.Golden.Value;
        var (_, tgtNames, _) = DlFixtures.GoldenNames();

        var graph = SameTarget.Build(DlFixtures.SboxRig.Value, DlFixtures.Stats.Value, align: false);
        Assert.Equal(tgtNames.Length, graph.NodeCount);

        var map = DlFixtures.MapNames(tgtNames, graph.NodeNames);
        var goldenLo = golden.Get("tgt.lo").Data;
        var goldenGo = golden.Get("tgt.go").Data;
        var goldenX = golden.Get("tgt.x").Data;

        // cm-scale geometry; the golden side round-tripped through a 6-decimal BVH.
        const float geomTolerance = 1e-2f;
        for (var g = 0; g < tgtNames.Length; g++)
        {
            var n = map[g];
            for (var c = 0; c < 3; c++)
            {
                Assert.True(MathF.Abs(goldenLo[g * 3 + c] - graph.LoRaw[n * 3 + c]) <= geomTolerance,
                    $"lo[{tgtNames[g]}][{c}]: golden {goldenLo[g * 3 + c]} vs {graph.LoRaw[n * 3 + c]}");
                Assert.True(MathF.Abs(goldenGo[g * 3 + c] - graph.GoRaw[n * 3 + c]) <= geomTolerance,
                    $"go[{tgtNames[g]}][{c}]: golden {goldenGo[g * 3 + c]} vs {graph.GoRaw[n * 3 + c]}");
            }
            for (var c = 0; c < SameModel.SkelDim; c++)
            {
                Assert.True(MathF.Abs(goldenX[g * 6 + c] - graph.X[n * 6 + c]) <= ModelTolerance,
                    $"tgt_x[{tgtNames[g]}][{c}]");
            }
        }
    }

    [Fact]
    public void DecoderMatchesGoldenOnPortBuiltTargetGraph()
    {
        var golden = DlFixtures.Golden.Value;
        var model = DlFixtures.Model.Value;
        var (_, tgtNames, frames) = DlFixtures.GoldenNames();

        var graph = SameTarget.Build(DlFixtures.SboxRig.Value, DlFixtures.Stats.Value, align: false);
        var (x, edgeSrc, edgeDst, batch) = SameTarget.Tile(graph, frames);
        var hatD = model.Decode(golden.Get("enc.z").Data, x, edgeSrc, edgeDst, batch);

        // The GAT is permutation-equivariant: rows must match per joint NAME even though
        // this port orders nodes differently from the Python BVH traversal.
        var map = DlFixtures.MapNames(tgtNames, graph.NodeNames);
        var goldenD = golden.Get("dec.hatD").Data;
        var worst = 0f;
        for (var f = 0; f < frames; f++)
        {
            for (var g = 0; g < tgtNames.Length; g++)
            {
                var ours = (f * graph.NodeCount + map[g]) * SameModel.OutDim;
                var theirs = (f * tgtNames.Length + g) * SameModel.OutDim;
                for (var c = 0; c < SameModel.OutDim; c++)
                    worst = MathF.Max(worst, MathF.Abs(goldenD[theirs + c] - hatD[ours + c]));
            }
        }
        Assert.True(worst <= 1e-3f, $"hatD deviates from golden by {worst} on the port-built graph.");
    }

    // ---------------------------------------------------------------- end-to-end vs golden

    [Fact]
    public void FullPipelineMatchesGoldenEmbeddings()
    {
        var golden = DlFixtures.Golden.Value;
        var model = DlFixtures.Model.Value;
        var (_, _, frames) = DlFixtures.GoldenNames();

        var scene = DlFixtures.FixtureBvhTruncated(frames + 2);
        var graph = SameFeatures.BuildSourceGraph(
            scene, 0, map: null, DlFixtures.Stats.Value,
            new SameFeatures.SourceOptions
            {
                NativeFrameDrop = true,
                Align = false,
                GroundShift = false,
            });
        var z = model.Encode(graph.X, graph.EdgeSrc, graph.EdgeDst, graph.Batch, graph.FrameCount);

        var goldenZ = golden.Get("enc.z").Data;
        var worst = 0f;
        for (var i = 0; i < goldenZ.Length; i++)
            worst = MathF.Max(worst, MathF.Abs(goldenZ[i] - z[i]));
        Assert.True(worst <= 1e-3f,
            $"End-to-end embedding deviates from golden by {worst} (importer + feature + model stack).");
    }

    private static void AssertClose(float[] expected, float[] actual, float tolerance, string what)
    {
        Assert.Equal(expected.Length, actual.Length);
        var worst = 0f;
        var at = -1;
        for (var i = 0; i < expected.Length; i++)
        {
            var d = MathF.Abs(expected[i] - actual[i]);
            if (d > worst)
            {
                worst = d;
                at = i;
            }
        }
        Assert.True(worst <= tolerance, $"{what} deviates by {worst} (> {tolerance}) at flat index {at}.");
    }
}