feat(train): WiFlow-STD PyTorch->tch weight import + numerical parity proof

export_to_safetensors.py maps the retrained checkpoint (295 tensors -> 248
mapped, param sum exactly 2,225,042; num_batches_tracked dropped) into a
tch-loadable safetensors plus a deterministic parity fixture. Gated #[ignore]
integration test loads it strictly and asserts forward-pass agreement:
max abs diff 1.192e-7 on the seed-42 fixture. dump_variable_names test makes
the tch name layout authoritative. Zero architecture discrepancies found.

Co-Authored-By: claude-flow <ruv@ruv.net>
This commit is contained in:
ruv
2026-06-10 20:23:38 -04:00
parent cc50e28a29
commit af6621005f
4 changed files with 307 additions and 2 deletions
+2
View File
@@ -13,4 +13,6 @@ downloads/
*.npz
*.zip
*.mat
*.safetensors
results/parity_fixture.json
__pycache__/
@@ -0,0 +1,187 @@
"""ADR-152 §2.2: export the retrained WiFlow-STD PyTorch checkpoint to
safetensors with tch-rs (VarStore) variable names, plus a numerical-parity
fixture for the Rust port.
Outputs (all under results/, gitignored):
retrained_wiflow_std.safetensors -- 248 f32 tensors named exactly as the
Rust WiFlowStdModel VarStore expects
(see wiflow_std/model.rs
`dump_variable_names` for the
authoritative name dump)
parity_fixture.npz -- deterministic input (seed 42,
shape (2, 540, 20), uniform [0,1]) and
the Python model's eval-mode output
parity_fixture.json -- same data as flattened f32 lists, for
the dependency-free Rust test
(tests/test_wiflow_std_parity.rs)
PyTorch -> tch key mapping (derived from the VarStore dump, not guessed):
tcn.network.{i}.conv1_group.weight -> tcn{i}.conv1_group.weight
tcn.network.{i}.bn*_{group,pw}.<leaf> -> tcn{i}.bn*_{group,pw}.<leaf>
tcn.network.{i}.downsample.0.weight -> tcn{i}.ds_conv.weight
tcn.network.{i}.downsample.1.<leaf> -> tcn{i}.ds_bn.<leaf>
up.block.{0,1,4,5,8,9}.<leaf> -> conv_in.{conv1,bn1,conv2,bn2,conv3,bn3}.<leaf>
up.downsample.{0,1}.<leaf> -> conv_in.{ds_conv,ds_bn}.<leaf>
residual_blocks.{i}.block.{...}.<leaf> -> conv{i}.{conv1..bn3}.<leaf>
residual_blocks.{i}.downsample.{0,1} -> conv{i}.{ds_conv,ds_bn}
attention.{width,height}_axis.qkv_transform.weight
-> attention.{width,height}.qkv.weight
attention.{width,height}_axis.bn_* -> attention.{width,height}.bn_*
decoder.{0,1,3,4}.<leaf> -> {dec_conv1,dec_bn1,dec_conv2,dec_bn2}.<leaf>
*.num_batches_tracked -> dropped (tch BatchNorm has no such buffer)
Legacy upstream names (att. -> attention., final_conv. -> decoder.) are
remapped first, exactly as eval_repro.py does for the released checkpoint.
Usage:
.venv/Scripts/python.exe export_to_safetensors.py
"""
import json
import os
import re
import sys
import numpy as np
import torch
from safetensors.torch import save_file
HERE = os.path.dirname(os.path.abspath(__file__))
UPSTREAM = os.path.join(HERE, "upstream")
RESULTS = os.path.join(HERE, "results")
sys.path.insert(0, UPSTREAM)
# Upstream models/__init__.py is broken as published (imports a name tcn.py
# does not define); register a stub package so it never executes.
import types # noqa: E402
_models_pkg = types.ModuleType("models")
_models_pkg.__path__ = [os.path.join(UPSTREAM, "models")]
sys.modules["models"] = _models_pkg
from models.pose_model import WiFlowPoseModel # noqa: E402
CHECKPOINT = os.path.join(RESULTS, "retrained_best_pose_model.pth")
# Sequential index -> tch sub-name inside one ConvBlock1/AsymmetricConvBlock:
# [Conv2d(0), BN(1), SiLU(2), Dropout2d(3), Conv2d(4), BN(5), SiLU(6),
# Dropout2d(7), Conv2d(8), BN(9)]
_BLOCK_IDX = {"0": "conv1", "1": "bn1", "4": "conv2", "5": "bn2",
"8": "conv3", "9": "bn3"}
_DS_IDX = {"0": "ds_conv", "1": "ds_bn"}
_DECODER_IDX = {"0": "dec_conv1", "1": "dec_bn1", "3": "dec_conv2",
"4": "dec_bn2"}
def _conv_block(new_prefix: str, rest: str) -> str:
m = re.fullmatch(r"block\.(\d+)\.(.+)", rest)
if m:
return f"{new_prefix}.{_BLOCK_IDX[m.group(1)]}.{m.group(2)}"
m = re.fullmatch(r"downsample\.(\d+)\.(.+)", rest)
if m:
return f"{new_prefix}.{_DS_IDX[m.group(1)]}.{m.group(2)}"
raise KeyError(f"unmapped conv-block key: {new_prefix} / {rest}")
def map_key(key: str) -> str:
"""Map one PyTorch state_dict key to the tch VarStore name."""
m = re.fullmatch(r"tcn\.network\.(\d+)\.(.+)", key)
if m:
i, rest = m.groups()
rest = (rest.replace("downsample.0.", "ds_conv.")
.replace("downsample.1.", "ds_bn."))
return f"tcn{i}.{rest}"
m = re.fullmatch(r"up\.(.+)", key)
if m:
return _conv_block("conv_in", m.group(1))
m = re.fullmatch(r"residual_blocks\.(\d+)\.(.+)", key)
if m:
return _conv_block(f"conv{m.group(1)}", m.group(2))
m = re.fullmatch(r"attention\.(width|height)_axis\.(.+)", key)
if m:
axis, rest = m.groups()
rest = rest.replace("qkv_transform.", "qkv.")
return f"attention.{axis}.{rest}"
m = re.fullmatch(r"decoder\.(\d+)\.(.+)", key)
if m:
return f"{_DECODER_IDX[m.group(1)]}.{m.group(2)}"
raise KeyError(f"unmapped checkpoint key: {key}")
def main():
state = torch.load(CHECKPOINT, map_location="cpu", weights_only=True)
if not isinstance(state, dict) or "tcn.network.0.conv1_group.weight" not in {
k for k in state
} | {k.replace("att.", "attention.") for k in state}:
# tolerate trainer wrappers like {"model_state_dict": ...}
for wrapper in ("model_state_dict", "state_dict", "model"):
if isinstance(state, dict) and wrapper in state:
state = state[wrapper]
break
# Legacy upstream names predate the published code (eval_repro.py).
renames = {"att.": "attention.", "final_conv.": "decoder."}
state = {next((new + k[len(old):] for old, new in renames.items()
if k.startswith(old)), k): v
for k, v in state.items()}
mapped = {}
dropped = 0
for k, v in state.items():
if k.endswith("num_batches_tracked"):
dropped += 1
continue
tch_key = map_key(k)
if tch_key in mapped:
raise KeyError(f"duplicate mapped key: {k} -> {tch_key}")
mapped[tch_key] = v.detach().to(torch.float32).contiguous()
n_params = sum(v.numel() for k, v in mapped.items()
if "running_" not in k)
print(f"checkpoint tensors: {len(state)} "
f"(dropped {dropped} num_batches_tracked)")
print(f"mapped tensors: {len(mapped)}, "
f"non-buffer params: {n_params/1e6:.6f}M")
assert len(mapped) == 248, f"expected 248 tch variables, got {len(mapped)}"
assert n_params == 2_225_042, f"param count mismatch: {n_params}"
st_path = os.path.join(RESULTS, "retrained_wiflow_std.safetensors")
save_file(mapped, st_path)
print(f"wrote {st_path}")
# ---- parity fixture --------------------------------------------------
model = WiFlowPoseModel(dropout=0.5)
model.load_state_dict(state, strict=True)
model.eval()
gen = torch.Generator().manual_seed(42)
x = torch.rand(2, 540, 20, generator=gen, dtype=torch.float32)
with torch.no_grad():
y = model(x)
print(f"fixture input {tuple(x.shape)} -> output {tuple(y.shape)}, "
f"output range [{y.min().item():.6f}, {y.max().item():.6f}]")
np.savez(os.path.join(RESULTS, "parity_fixture.npz"),
input=x.numpy(), output=y.numpy())
fixture = {
"seed": 42,
"input_shape": list(x.shape),
"input": x.flatten().tolist(),
"output_shape": list(y.shape),
"output": y.flatten().tolist(),
}
json_path = os.path.join(RESULTS, "parity_fixture.json")
with open(json_path, "w") as f:
json.dump(fixture, f)
print(f"wrote {os.path.join(RESULTS, 'parity_fixture.npz')}")
print(f"wrote {json_path}")
if __name__ == "__main__":
main()
@@ -3,8 +3,12 @@
//! Idiomatic reimplementation of the DY2434 reference (Apache-2.0); see the
//! [module docs](crate::wiflow_std) for provenance and the evidence grade.
//! Weights are initialised from scratch (tch defaults; the axial-attention
//! qkv conv mirrors the reference's `N(0, sqrt(1/in_planes))` init). Loading
//! the retrained PyTorch checkpoint is a follow-up (key remap + `vs.load`).
//! qkv conv mirrors the reference's `N(0, sqrt(1/in_planes))` init). The
//! retrained PyTorch checkpoint loads via [`WiFlowStdModel::load`] after
//! key-remapped safetensors export
//! (`benchmarks/wiflow-std/export_to_safetensors.py`); numerical parity with
//! the PyTorch forward pass is proven by
//! `tests/test_wiflow_std_parity.rs` (max abs diff ~1.2e-7).
use tch::{nn, Device, Tensor};
@@ -267,6 +271,25 @@ mod tests {
);
}
/// Dumps the authoritative tch `VarStore` variable names + shapes. This is
/// the source of truth for the PyTorch→tch key mapping implemented in
/// `benchmarks/wiflow-std/export_to_safetensors.py` — rerun it (with
/// `--nocapture`) whenever the architecture changes.
#[test]
fn dump_variable_names() {
let cfg = WiFlowStdConfig::default();
let model = WiFlowStdModel::new(&cfg, Device::Cpu).expect("build");
let vars = model.var_store().variables();
let mut names: Vec<(String, Vec<i64>)> =
vars.iter().map(|(n, t)| (n.clone(), t.size())).collect();
names.sort();
for (name, shape) in &names {
println!("{name} {shape:?}");
}
println!("total: {} variables", names.len());
assert!(!names.is_empty());
}
#[test]
fn invalid_config_is_rejected() {
let cfg = WiFlowStdConfig {
@@ -0,0 +1,93 @@
//! Numerical parity between the Rust WiFlow-STD port and the retrained
//! PyTorch checkpoint (ADR-152 §2.2).
//!
//! The fixtures are produced by `benchmarks/wiflow-std/export_to_safetensors.py`
//! (gitignored — they derive from the retrained checkpoint, which is itself
//! gitignored):
//!
//! - `results/retrained_wiflow_std.safetensors` — the epoch-36 checkpoint
//! (val PCK@20 96.99%) remapped to tch `VarStore` variable names
//! - `results/parity_fixture.json` — a deterministic input (seed 42, shape
//! `(2, 540, 20)`, uniform `[0, 1]`) and the upstream `WiFlowPoseModel`'s
//! eval-mode output on it
//!
//! Run explicitly (needs LibTorch, e.g. `LIBTORCH_USE_PYTORCH=1` with the
//! torch DLL directory on `PATH`):
//!
//! ```text
//! cargo test -p wifi-densepose-train --features tch-backend \
//! --test test_wiflow_std_parity -- --ignored --nocapture
//! ```
#![cfg(feature = "tch-backend")]
use std::fs::File;
use std::io::BufReader;
use std::path::PathBuf;
use tch::{Device, Tensor};
use wifi_densepose_train::{WiFlowStdConfig, WiFlowStdModel};
#[derive(serde::Deserialize)]
struct ParityFixture {
input_shape: Vec<i64>,
input: Vec<f32>,
output_shape: Vec<i64>,
output: Vec<f32>,
}
fn results_dir() -> PathBuf {
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("..")
.join("..")
.join("..")
.join("benchmarks")
.join("wiflow-std")
.join("results")
}
/// Loads the retrained checkpoint into the Rust model and asserts the forward
/// pass matches PyTorch to within 1e-4 max absolute difference.
///
/// `#[ignore]`d by default: it needs the gitignored fixtures above plus a
/// working LibTorch environment, neither of which exist in CI.
#[test]
#[ignore = "needs gitignored fixtures (run export_to_safetensors.py) + LibTorch env; run with --ignored"]
fn retrained_checkpoint_matches_pytorch_forward() {
let dir = results_dir();
let weights = dir.join("retrained_wiflow_std.safetensors");
let fixture_path = dir.join("parity_fixture.json");
for p in [&weights, &fixture_path] {
assert!(
p.exists(),
"missing fixture {} — run benchmarks/wiflow-std/export_to_safetensors.py first",
p.display()
);
}
let fixture: ParityFixture = serde_json::from_reader(BufReader::new(
File::open(&fixture_path).expect("open parity_fixture.json"),
))
.expect("parse parity_fixture.json");
assert_eq!(fixture.input_shape, vec![2, 540, 20]);
assert_eq!(fixture.output_shape, vec![2, 15, 2]);
let cfg = WiFlowStdConfig::default();
let mut model = WiFlowStdModel::new(&cfg, Device::Cpu).expect("build default model");
model
.load(&weights)
.expect("safetensors load: every VarStore variable must match by name and shape");
let input = Tensor::from_slice(&fixture.input).reshape(&fixture.input_shape[..]);
let expected = Tensor::from_slice(&fixture.output).reshape(&fixture.output_shape[..]);
let output = model.forward_inference(&input);
assert_eq!(output.size(), fixture.output_shape);
let max_diff = (&output - &expected).abs().max().double_value(&[]);
println!("max |rust - python| = {max_diff:.3e}");
assert!(
max_diff < 1e-4,
"Rust forward pass diverges from PyTorch: max abs diff {max_diff:.3e} >= 1e-4"
);
}