mirror of
https://github.com/ruvnet/RuView
synced 2026-07-30 18:41:42 +00:00
17471e93ff
* feat(calibration): NodeGeometry transceiver-geometry recording (ADR-152 §2.1.1) PerceptAlign-motivated geometry capture at enrollment: per-node optional records (position, antenna orientation, inter-node distances, acquisition method) — recorded when known, never required. Event-sourced via EnrollmentEvent::GeometryRecorded (latest recording wins); persisted on SpecialistBank with serde defaults so pre-ADR-152 bank JSON loads cleanly (fixture-proven, and geometry-free banks serialize byte-shape-identical to the old schema); threaded through MultiNodeMixture as data only — the learned geometry embeddings and algorithmic fusion use are §2.1.2, deliberately deferred until the ADR-151 P6 LoRA heads exist. Geometry recorded from now on means banks captured today remain usable for layout-conditioned training later — you can't retroactively add geometry to data you didn't record. 8 new tests (3 geometry, 2 anchor, 2 bank, 1 multistatic) + full-loop extension (2-node geometry, one tape-measured + one unknown, surviving the bank JSON round-trip the runtime loads from). 50/50 calibration (both feature configs) + 23 CLI tests green. Co-Authored-By: RuFlo <ruv@ruv.net> * feat(training): two-checkerboard camera↔room calibration for ADR-079 labels (ADR-152 §2.1.3) Defends the camera-supervised pipeline against PerceptAlign's "coordinate overfitting": MediaPipe keypoints were emitted in raw camera coordinates with no shared frame and no transceiver-geometry metadata — the exact label shape that memorizes deployment layout and collapses cross-layout. - scripts/calibrate-camera-room.py + calibration_lib.py: OpenCV two-checkerboard calibration → versioned bundle JSON (intrinsics, camera→room extrinsics, checkerboard spec, transceiver geometry, sha256 calibration_id). Intrinsics resolve from file > cache > multi-view computation > loud-warning 2-view fallback. - collect-ground-truth.py --calibration <bundle>: every sample gains keypoints_room (unit bearing rays from the camera center in the room frame — documented projective alignment; raw image coords preserved so training chooses), camera_origin_room, calibration_id, and the transceiver geometry stamp. Without the flag, output is byte-identical to before (tested) + a one-line ADR-152 warning. Design finding (recorded for ADR-152): a single planar checkerboard's corner grid is centrosymmetric — the reversed corner ordering fits a ghost camera pose with IDENTICAL reprojection error, so per-board flip disambiguation is mathematically ill-posed. solve_two_board_extrinsics solves the joint wall+floor set over all 4 flip combinations, where the minimum is unique — an independent reason the TWO-checkerboard method is required, beyond what PerceptAlign states. 15 headless pytest tests green (synthetic corners: extrinsics recovery incl. ghost resolution, bundle round-trip + hash stability, ray transforms w/ distortion + cross-resolution, no-calibration byte identity). Co-Authored-By: RuFlo <ruv@ruv.net> * feat(benchmarks): WiFlow-STD reproduction harness + measurement (a) results (ADR-152 §2.2) Shipped checkpoint REFUTED (0.08% PCK@20, wrong keypoint normalization); 6 reproducibility defects documented (broken imports, corrupted dataset tail with float32-max garbage that NaN-poisons fp16 BatchNorm, unreachable test phase). After repairs, retraining with upstream defaults reproduces 96.09% PCK@20 full-test / 96.61% corruption-free (published 97.25%) on RTX 5080. Claims graded MEASURED-EQUIVALENT; 2.23M params + ~0.055 GFLOPs verified. Third-party code/weights/data stay out of tree (gitignored). Co-Authored-By: claude-flow <ruv@ruv.net> * feat: ADR-152 Rust integrations + ADR-153 802.11bf protocol model - calibration: GeometryEmbedding — 32-slot permutation-invariant NodeGeometry featurization for future LoRA-head conditioning (ADR-152 §2.1.2); derived SpecialistBank::geometry_embedding() accessor; 59 tests - train: MaePretrainConfig + patchify/random-mask with UNSW measured recipe (80% masking, (30,3) patches; ADR-152 §2.3, arXiv 2511.18792); strict no-truncate/no-NaN policy; proptest properties - train: WiFlowStdModel — tch-gated port of the verified ~96%-PCK@20 WiFlow-STD architecture (ADR-152 §2.2 beyond-SOTA); ungated param formula pinned to 2,225,042; 15/17-keypoint support; 239 crate tests - hardware: ieee80211bf forward-compatibility protocol model (ADR-153): SpecProfile gates, SensingCapabilities negotiation, required ConsentMode, session FSM, SensingTransport + SimTransport + OpportunisticCsiBridge; full acceptance checklist covered; 156+4 tests - deps: ruvector bumps per ADR-152 §2.6 survey (mincut/solver 2.0.6, attention 2.1.0, gnn 2.2.0); vendor/ruvector synced to a083bd77f - docs: ADR-153 accepted; ADR-152 §2.2 status, §2.4 amendment, §2.6 added Workspace: 162 test suites green (--no-default-features); Python proof PASS. Known pre-existing flake: homecore-api env_empty_falls_back_to_defaults (unserialized env-var mutation) — untouched, follow-up. Co-Authored-By: claude-flow <ruv@ruv.net> * docs: CHANGELOG + CLAUDE.md entries for ADR-152 integrations and ADR-153 Co-Authored-By: claude-flow <ruv@ruv.net> * fix(train): repair tch-backend bit-rot — gated path compiles and tests run again Mechanical API refresh against current tch: Vec::from(Tensor) -> try_from (+ explicit flatten), numel() usize cast, Rem/div ops -> remainder() / divide_scalar_mode(floor) — the latter fixed a silent true-division bug in heatmap argmax decoding; clamp(1.0, f64::MAX) -> clamp_min (torch 2.x scalar overflow panic); petgraph EdgeRef import; missing EvalMetrics and verify_checkpoint_dir APIs that tests documented. wiflow_std roundtrip test uses safetensors (.pt _save_parameters roundtrip broken in torch 2.11 Windows). Gated: 349 passed (incl. all 20 wiflow_std); ungated: unchanged. Known pre-existing: gaussian-heatmap convention mismatch (2 tests), proof seed race under parallel threads — documented, deliberate follow-ups. Co-Authored-By: claude-flow <ruv@ruv.net> * 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> * fix: workflow-review findings — BN gamma init, ThresholdParams serde, init docs Concurrent validation workflow (2 review lanes + adversarial verification, 13 agents): 5 confirmed findings, 3 refuted. Fixes: - wiflow_std: pin BatchNorm gamma to 1.0 (tch default draws Uniform(0,1) — silently halves activations in from-scratch training; loaded checkpoints unaffected, parity re-verified after the change) - wiflow_std: document the conv-init divergences vs the reference's effective kaiming_normal(fan_out) re-init (from-scratch dynamics only) - ieee80211bf: ThresholdParams deserialization validates via try_from so the <=100 invariant holds for untrusted payloads (+ rejection test) Benchmarks (release, ruvzen): GeometryEmbedding 1.84us/call (542k/s), MAE tokenization 7.38us/window (135k/s), 802.11bf FSM 8.9M events/s — nothing suspicious. Co-Authored-By: claude-flow <ruv@ruv.net> * docs(adr): ADR-152 §2.1.4 gate resolved — PerceptAlign repo MIT, dataset on HF Co-Authored-By: claude-flow <ruv@ruv.net> * feat(benchmarks): edge optimization measured + measurement (b) blocked + 92.9% retraction Edge optimization (ADR-152 optimize track): ONNX Runtime fp32 is the CPU latency win (3.2 ms/window, ~3.4x faster than torch, parity 2.4e-7); ORT dynamic int8 reaches 2.44 MB (paper's ~2.2 MB claim plausible only via conv-capable toolchains; -0.16pt PCK@20, +18% MPJPE, 2x slower); torch dynamic quant converts 0% of this conv-only model; fp16 halves storage free but is slower on CPU. Measurement (b) BLOCKED-ON-DATA: only 1,077 paired ESP32 windows exist (stop rule <2k). Forensic recheck of the surviving April holdout RETRACTS the ADR-079 '92.9% PCK@20' figure: constant-output model, absolute (not torso) threshold, 69 near-static frames — mean predictor scores 100% under that protocol; torso-PCK@20 is 19.1%. Corroborates PR #535. Stale citations removed from user-guide, readme-details, ADR-152 §2.1.3; no-citation rule extended to ADR-079 accuracy claims. Unblock: >=2k-window multi-pose paired session + torso-PCK re-baseline. Co-Authored-By: claude-flow <ruv@ruv.net> * docs(user-guide): corrected camera-supervised collection tutorial Step 0 CSI-rate check + session-length math (window yield = frames/20 — the May session's 8x under-delivery was a ~12 Hz CSI rate, not an aligner bug); two-checkerboard calibration step (ADR-152 §2.1.3); pose-variety and confidence guidance; torso-normalized PCK + temporal-split + pred-variance eval protocol (lessons from the 92.9% retraction); scale presets re-keyed to realistic window counts. Co-Authored-By: claude-flow <ruv@ruv.net> * feat(benchmarks): static PTQ int8 (calibrated) results + overnight capture script Conv-only static QDQ beats dynamic int8 on accuracy (PCK@20 96.61-96.63% vs 96.52%, MPJPE +10% vs +18% over fp32) at ~equal size/latency; all-ops QDQ strictly worse (int8 activations through attention glue). Entropy calibration verified bit-identical to MinMax on this data. Deployment: ONNX fp32 for speed (3.2ms), static conv-only QDQ for smallest (2.53MB). Also: scripts/overnight-empty-capture.py — segmented UDP CSI recorder for empty-room baselines (no glob collisions, detach-safe). Co-Authored-By: claude-flow <ruv@ruv.net> * feat(benchmarks): measurement (b) MEASURED — optimization transfer only, mean-pose baseline wins WiFlow-STD fine-tuned on 2,046 fresh single-room ESP32 paired windows (temporal 70/15/15, 70->540 adapter, K=17): pretrained-init 65% PCK@20 vs scratch 0% (optimization transfer) but frozen-trunk ~0% (no feature transfer), and NOTHING beats the mean-pose baseline (95.9% PCK@20 — single subject, near-static normalized coords). Honesty gates held: pred std 0.0113 (non-constant model) but mean-baseline dominance means no citable CSI->pose capability from this data. ADR-152 open question 1 answered partially; definitive answer needs multi-subject/position data. Two new aligner findings: heterogeneous csi_shape with silent zero-padding (~20%), and extractCsiMatrix's transposed shape label (frame-major data, [nSc, nFrames] label) — fixes pending. Co-Authored-By: claude-flow <ruv@ruv.net> * feat(benchmarks): efficiency sweep MEASURED — half model dominates full reference Compact WiFlow-STD variants on the same data/split/protocol: half (843,834 params, 0.38x) strictly dominates the 2.23M reference (PCK@20 96.62 vs 96.61, PCK@50 99.47 vs 99.11, MPJPE 0.00898 vs 0.0094) — the published architecture is over-parameterized for its own benchmark. quarter (338k) 96.05%; tiny (56,290 params, 1/39.5) holds 94.11% — a ~220KB fp32 edge candidate. In-domain caveats recorded; cross-domain untested. Co-Authored-By: claude-flow <ruv@ruv.net> * feat(train): compact WiFlow-STD presets in Rust + tiny edge artifact (ADR-152) WiFlowStdConfig gains half()/quarter()/tiny() mirroring the overnight sweep exactly: TcnGroupsMode (Fixed/Gcd/Depthwise), input_pw_groups, derived stride schedule and decoder-mid (all default to upstream behavior; legacy serde JSON unaffected). Param formulas pin to trained ground truth first try: 843,834 / 338,600 / 56,290; default 2,225,042 pin and 1.192e-7 parity unchanged. 248 tests green. Tiny edge artifact (tiny_edge_bench.py): ONNX fp32 = 295 KB, 0.66 ms/win (~1,500/s CPU), 94.11% PCK@20 (matches sweep clean-test exactly; parity 1.49e-7). Static int8 is a bad trade at this scale (-1.43pt, +19% MPJPE, -16% size, slower) — recorded as negative result. Export note: width-16 breaks AdaptiveAvgPool((15,1)) TorchScript export; replaced by exact mean+matmul equivalent, proven by parity. Co-Authored-By: claude-flow <ruv@ruv.net> * fix: resolve all 10 confirmed code-review findings (7-angle review, 20/20 verified) wiflow_std: min_feature_width (default 15) replaces the keypoints->stride coupling — for_keypoints(17) now provably builds the trained [2,2,2,2] graph and pools 15->17, matching the validated Python protocol (pinned by tests); param_count() total on invalid configs; random_mask returns Result and rejects non-finite/out-of-range ratios; trainer checkpoints switched to safetensors (.pt VarStore roundtrip broken on Windows torch 2.11). ieee80211bf: SBP proxy now re-triggers instances and relays reports via Action::RelaySbpReport -> SensingFrame::SbpReport (clients consume via their existing path); missed_instances reset on success = consecutive semantics; SessionTable gains a guarded SBP entry point + unknown-id drop counter; initiator-role sessions reject inbound setup/SBP requests (RejectedNotSupported) closing the idle hijack; StartSetup/StartSbp outside Idle return InvalidStateForCommand; SBP validation unified through evaluate_setup with a 1:1 SetupStatus->SbpStatus mapping. events.rs split out to honor the 500-line cap. calibration/cli: enrollment geometry now actually reaches trained banks — both production call sites attach .with_geometry; --geometry flag on train-room and POST /enroll/geometry + train-body geometry on calibrate-serve give production a recording surface; geometry-free banks log the ADR-152 §2.1.2 note. benchmarks: corruption masks committed as ground truth (unregenerable after in-place cleaning; verified bit-identical regeneration from the pristine copy) + generate_corruption_masks.py producer; _bench_common.py dedups the 5x-copied shim/evaluate/seed/remap (post-refactor PCK@20 re-verified equal to the last digit); remote scripts get the mmap patch; tiny_edge --calib validated multiple-of-64; onnx_bench --help no longer executes (and overwrote) the export — artifact restored byte-exact. Workspace: 2,963 tests passed, 0 failed; Python proof PASS. Co-Authored-By: claude-flow <ruv@ruv.net> * ci: build workspace tests without debuginfo — runner disk exhaustion The combined 38-crate debug target exceeds the GitHub runner's disk ('final link failed: No space left on device'); the same tree measured 151GB locally with full debuginfo. CARGO_PROFILE_{DEV,TEST}_DEBUG=0 shrinks the target ~5-10x; debuginfo serves no purpose in CI test runs. Co-Authored-By: claude-flow <ruv@ruv.net>
397 lines
15 KiB
Rust
397 lines
15 KiB
Rust
//! Masked-autoencoder (MAE) pretraining recipe for the ADR-150 RF foundation
|
||
//! encoder — ADR-152 §2.3 (amends ADR-150 §2.3).
|
||
//!
|
||
//! Implements the *measured* tokenization recipe from the UNSW MAE pretraining
|
||
//! study (arXiv [2511.18792](https://arxiv.org/abs/2511.18792), Nov 2025), the
|
||
//! largest heterogeneous CSI pretraining run to date (1,320,892 samples, 14
|
||
//! public datasets, 4 devices, 2.4/5/6 GHz, 20–160 MHz):
|
||
//!
|
||
//! - **80% masking ratio** over the patch grid.
|
||
//! - **Small (30, 3) patches** — 30 time steps × 3 subcarriers — measured
|
||
//! **+4.7%** over (40, 5) patches by preserving fine temporal dynamics.
|
||
//! - Encoder capacity stays **ViT-Small-class (~15M params)**: ViT-Base adds
|
||
//! only +0.4–0.9% over ViT-Small in-study, corroborating ADR-150's own
|
||
//! finding that capacity hurts cross-subject transfer.
|
||
//! - Unseen-domain performance scales **log-linearly with pretraining data,
|
||
//! unsaturated at 1.3M samples** — data aggregation outranks architecture
|
||
//! work (ADR-152 §2.3).
|
||
//!
|
||
//! This module provides the GPU-free half of the recipe: configuration,
|
||
//! patchification, and deterministic random masking. The (future, ADR-150)
|
||
//! encoder consumes [`PatchGrid`] + [`MaskIndices`] to compute the masked
|
||
//! reconstruction loss (`L_masked_csi` in ADR-150 §2.3's loss stack).
|
||
//!
|
||
//! ## Axis convention
|
||
//!
|
||
//! A CSI window is `time × subcarriers`, row-major (`index = t * subc + sc`),
|
||
//! matching the crate's `[T, …, n_sc]` dataset layout (time first, subcarriers
|
||
//! last) and the UNSW "(30 time steps, 3 subcarriers)" patch framing. Patches
|
||
//! are indexed row-major over the patch grid (`p = pt * n_patches_subc + ps`),
|
||
//! and values within a patch are row-major time-major
|
||
//! (`local = lt * patch_subc + lsc`).
|
||
//!
|
||
//! ## Divisibility policy: error, never truncate
|
||
//!
|
||
//! Window dimensions **must** be exact multiples of the patch dimensions.
|
||
//! Non-divisible shapes return [`MaeError::NotDivisible`] instead of silently
|
||
//! truncating trailing samples (this crate never silently drops data). The
|
||
//! error names the largest divisible crop; use
|
||
//! [`MaePretrainConfig::cropped_window_shape`] to compute it and crop
|
||
//! explicitly before calling [`patchify`].
|
||
//!
|
||
//! ## Example
|
||
//!
|
||
//! ```rust
|
||
//! use wifi_densepose_train::mae::MaePretrainConfig;
|
||
//!
|
||
//! let cfg = MaePretrainConfig::default(); // 0.80 masking, (30, 3) patches
|
||
//! cfg.validate().expect("default recipe is valid");
|
||
//!
|
||
//! // 90 frames × 54 subcarriers → a 3 × 18 grid of (30, 3) patches.
|
||
//! let window = vec![0.25_f32; 90 * 54];
|
||
//! let (grid, mask) = cfg.mask_window(&window, 90, 54).unwrap();
|
||
//! assert_eq!(grid.n_patches(), 54);
|
||
//! assert_eq!(mask.masked.len(), 43); // round(0.80 * 54)
|
||
//! assert_eq!(mask.visible.len(), 11);
|
||
//! ```
|
||
|
||
use serde::{Deserialize, Serialize};
|
||
|
||
use crate::error::{ConfigError, MaeError};
|
||
use crate::virtual_aug::Xorshift64;
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// MaePretrainConfig
|
||
// ---------------------------------------------------------------------------
|
||
|
||
/// Hyper-parameters for masked-CSI pretraining (ADR-152 §2.3).
|
||
///
|
||
/// Defaults are the measured-optimal UNSW recipe (arXiv 2511.18792); change
|
||
/// them only with benchmark evidence. Serializable so the recipe is recorded
|
||
/// in checkpoint metadata alongside [`crate::config::TrainingConfig`].
|
||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||
pub struct MaePretrainConfig {
|
||
/// Fraction of patches hidden from the encoder, in `(0, 1)`.
|
||
///
|
||
/// Default: **0.80** (UNSW measured optimum).
|
||
pub mask_ratio: f64,
|
||
|
||
/// Patch extent along the time axis, in frames. Default: **30**.
|
||
pub patch_time: usize,
|
||
|
||
/// Patch extent along the subcarrier axis. Default: **3**.
|
||
pub patch_subc: usize,
|
||
|
||
/// Base seed for the deterministic mask sampler. Default: **42**.
|
||
///
|
||
/// For per-sample masks derive a child seed (e.g.
|
||
/// `seed ^ sample_idx as u64`) and pass it to [`random_mask`]; reusing one
|
||
/// seed yields the identical mask for every sample.
|
||
pub seed: u64,
|
||
}
|
||
|
||
impl Default for MaePretrainConfig {
|
||
fn default() -> Self {
|
||
MaePretrainConfig {
|
||
mask_ratio: 0.80,
|
||
patch_time: 30,
|
||
patch_subc: 3,
|
||
seed: 42,
|
||
}
|
||
}
|
||
}
|
||
|
||
impl MaePretrainConfig {
|
||
/// Validate the shape-independent fields.
|
||
///
|
||
/// # Validated invariants
|
||
///
|
||
/// - `mask_ratio` must be strictly inside `(0, 1)` and finite.
|
||
/// - `patch_time` and `patch_subc` must be at least 1.
|
||
pub fn validate(&self) -> Result<(), ConfigError> {
|
||
if !self.mask_ratio.is_finite() || self.mask_ratio <= 0.0 || self.mask_ratio >= 1.0 {
|
||
return Err(ConfigError::invalid_value(
|
||
"mask_ratio",
|
||
format!("must be in (0.0, 1.0), got {}", self.mask_ratio),
|
||
));
|
||
}
|
||
if self.patch_time == 0 {
|
||
return Err(ConfigError::invalid_value("patch_time", "must be >= 1"));
|
||
}
|
||
if self.patch_subc == 0 {
|
||
return Err(ConfigError::invalid_value("patch_subc", "must be >= 1"));
|
||
}
|
||
Ok(())
|
||
}
|
||
|
||
/// Check this recipe against a concrete `time × subc` window shape.
|
||
///
|
||
/// Errors if a patch dimension exceeds the window or if either axis is
|
||
/// not an exact multiple of the patch extent (divisibility policy above).
|
||
pub fn validate_for_window(&self, time: usize, subc: usize) -> Result<(), MaeError> {
|
||
check_axis("time", time, self.patch_time)?;
|
||
check_axis("subcarrier", subc, self.patch_subc)?;
|
||
Ok(())
|
||
}
|
||
|
||
/// Largest `(time, subc)` crop of the given window that is exactly
|
||
/// divisible by the patch dimensions. Either component may be 0 when the
|
||
/// window is smaller than one patch.
|
||
#[must_use]
|
||
pub fn cropped_window_shape(&self, time: usize, subc: usize) -> (usize, usize) {
|
||
(
|
||
(time / self.patch_time) * self.patch_time,
|
||
(subc / self.patch_subc) * self.patch_subc,
|
||
)
|
||
}
|
||
|
||
/// Number of patches a `time × subc` window yields under this recipe.
|
||
pub fn num_patches(&self, time: usize, subc: usize) -> Result<usize, MaeError> {
|
||
self.validate_for_window(time, subc)?;
|
||
Ok((time / self.patch_time) * (subc / self.patch_subc))
|
||
}
|
||
|
||
/// Exact number of masked patches for a grid of `n_patches`:
|
||
/// `round(mask_ratio * n_patches)`, clamped to `[0, n_patches]`.
|
||
#[must_use]
|
||
pub fn num_masked(&self, n_patches: usize) -> usize {
|
||
((self.mask_ratio * n_patches as f64).round() as usize).min(n_patches)
|
||
}
|
||
|
||
/// Patchify `window` and draw the deterministic random mask in one step,
|
||
/// using `self.seed`. See [`patchify`] and [`random_mask`].
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// Everything [`patchify`] rejects, plus [`MaeError::InvalidMaskRatio`]
|
||
/// if `self.mask_ratio` is not finite or outside `(0, 1)` (the
|
||
/// [`Self::validate`] rule) — a NaN ratio must never silently mask zero
|
||
/// patches.
|
||
pub fn mask_window(
|
||
&self,
|
||
window: &[f32],
|
||
time: usize,
|
||
subc: usize,
|
||
) -> Result<(PatchGrid, MaskIndices), MaeError> {
|
||
let grid = patchify(window, time, subc, self)?;
|
||
let mask = random_mask(grid.n_patches(), self.mask_ratio, self.seed)?;
|
||
Ok((grid, mask))
|
||
}
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// PatchGrid / MaskIndices
|
||
// ---------------------------------------------------------------------------
|
||
|
||
/// A CSI window decomposed into non-overlapping `patch_time × patch_subc`
|
||
/// patches (see the module-level axis convention).
|
||
#[derive(Debug, Clone, PartialEq)]
|
||
pub struct PatchGrid {
|
||
/// Patch extent along the time axis.
|
||
pub patch_time: usize,
|
||
/// Patch extent along the subcarrier axis.
|
||
pub patch_subc: usize,
|
||
/// Number of patch rows (`time / patch_time`).
|
||
pub n_patches_time: usize,
|
||
/// Number of patch columns (`subc / patch_subc`).
|
||
pub n_patches_subc: usize,
|
||
/// Flattened patches, row-major over the grid; each inner `Vec` is one
|
||
/// patch of length `patch_time * patch_subc`, row-major time-major.
|
||
pub patches: Vec<Vec<f32>>,
|
||
}
|
||
|
||
impl PatchGrid {
|
||
/// Total number of patches in the grid.
|
||
#[must_use]
|
||
pub fn n_patches(&self) -> usize {
|
||
self.n_patches_time * self.n_patches_subc
|
||
}
|
||
|
||
/// Number of scalar values per patch.
|
||
#[must_use]
|
||
pub fn patch_len(&self) -> usize {
|
||
self.patch_time * self.patch_subc
|
||
}
|
||
|
||
/// Window shape `(time, subc)` this grid reconstructs to.
|
||
#[must_use]
|
||
pub fn window_shape(&self) -> (usize, usize) {
|
||
(
|
||
self.n_patches_time * self.patch_time,
|
||
self.n_patches_subc * self.patch_subc,
|
||
)
|
||
}
|
||
}
|
||
|
||
/// Sorted, disjoint patch-index sets produced by [`random_mask`]. Together
|
||
/// they cover `0..n_patches` exactly.
|
||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||
pub struct MaskIndices {
|
||
/// Indices of patches hidden from the encoder (`round(ratio * n)` of them).
|
||
pub masked: Vec<usize>,
|
||
/// Indices of patches the encoder sees.
|
||
pub visible: Vec<usize>,
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// patchify / unpatchify
|
||
// ---------------------------------------------------------------------------
|
||
|
||
/// Decompose a row-major `time × subc` CSI window into the patch grid defined
|
||
/// by `cfg`.
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// - [`MaeError::WindowShapeMismatch`] if `window.len() != time * subc`.
|
||
/// - [`MaeError::PatchExceedsWindow`] / [`MaeError::NotDivisible`] per the
|
||
/// module-level divisibility policy.
|
||
/// - [`MaeError::NonFiniteValue`] on the first NaN/±inf encountered —
|
||
/// corrupted CSI must be cleaned upstream, never masked over (cf. the
|
||
/// WiFlow-STD NaN-poisoning incident, ADR-152 §2.2).
|
||
pub fn patchify(
|
||
window: &[f32],
|
||
time: usize,
|
||
subc: usize,
|
||
cfg: &MaePretrainConfig,
|
||
) -> Result<PatchGrid, MaeError> {
|
||
let expected = time * subc;
|
||
if window.len() != expected {
|
||
return Err(MaeError::WindowShapeMismatch {
|
||
time,
|
||
subc,
|
||
expected,
|
||
actual: window.len(),
|
||
});
|
||
}
|
||
cfg.validate_for_window(time, subc)?;
|
||
if let Some(idx) = window.iter().position(|v| !v.is_finite()) {
|
||
return Err(MaeError::NonFiniteValue {
|
||
row: idx / subc,
|
||
col: idx % subc,
|
||
value: window[idx],
|
||
});
|
||
}
|
||
|
||
let n_patches_time = time / cfg.patch_time;
|
||
let n_patches_subc = subc / cfg.patch_subc;
|
||
let mut patches = Vec::with_capacity(n_patches_time * n_patches_subc);
|
||
for pt in 0..n_patches_time {
|
||
for ps in 0..n_patches_subc {
|
||
let mut patch = Vec::with_capacity(cfg.patch_time * cfg.patch_subc);
|
||
for lt in 0..cfg.patch_time {
|
||
let t = pt * cfg.patch_time + lt;
|
||
let row_start = t * subc + ps * cfg.patch_subc;
|
||
patch.extend_from_slice(&window[row_start..row_start + cfg.patch_subc]);
|
||
}
|
||
patches.push(patch);
|
||
}
|
||
}
|
||
|
||
Ok(PatchGrid {
|
||
patch_time: cfg.patch_time,
|
||
patch_subc: cfg.patch_subc,
|
||
n_patches_time,
|
||
n_patches_subc,
|
||
patches,
|
||
})
|
||
}
|
||
|
||
/// Reassemble the full row-major `time × subc` window from a [`PatchGrid`].
|
||
/// Exact inverse of [`patchify`].
|
||
#[must_use]
|
||
pub fn unpatchify(grid: &PatchGrid) -> Vec<f32> {
|
||
unpatchify_select(grid, None, 0.0)
|
||
}
|
||
|
||
/// Reassemble the window keeping only the patches listed in `visible`;
|
||
/// every other patch's region is filled with `fill` (the standard MAE
|
||
/// "visible tokens + mask token" view of the input).
|
||
#[must_use]
|
||
pub fn unpatchify_visible(grid: &PatchGrid, visible: &[usize], fill: f32) -> Vec<f32> {
|
||
unpatchify_select(grid, Some(visible), fill)
|
||
}
|
||
|
||
fn unpatchify_select(grid: &PatchGrid, keep: Option<&[usize]>, fill: f32) -> Vec<f32> {
|
||
let (time, subc) = grid.window_shape();
|
||
let mut window = vec![fill; time * subc];
|
||
for (p, patch) in grid.patches.iter().enumerate() {
|
||
if let Some(keep) = keep {
|
||
if !keep.contains(&p) {
|
||
continue;
|
||
}
|
||
}
|
||
let pt = p / grid.n_patches_subc;
|
||
let ps = p % grid.n_patches_subc;
|
||
for lt in 0..grid.patch_time {
|
||
let t = pt * grid.patch_time + lt;
|
||
let row_start = t * subc + ps * grid.patch_subc;
|
||
let local_start = lt * grid.patch_subc;
|
||
window[row_start..row_start + grid.patch_subc]
|
||
.copy_from_slice(&patch[local_start..local_start + grid.patch_subc]);
|
||
}
|
||
}
|
||
window
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// random_mask
|
||
// ---------------------------------------------------------------------------
|
||
|
||
/// Draw a deterministic random mask over `n_patches` patches.
|
||
///
|
||
/// Exactly `round(mask_ratio * n_patches)` patches (clamped to
|
||
/// `[0, n_patches]`) are masked, chosen by a seeded Fisher–Yates shuffle
|
||
/// ([`Xorshift64`]), so the same `(n_patches, mask_ratio, seed)` triple always
|
||
/// yields the same mask. Both index lists are sorted ascending, disjoint, and
|
||
/// together cover `0..n_patches`.
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// [`MaeError::InvalidMaskRatio`] if `mask_ratio` is not finite or outside
|
||
/// the open interval `(0, 1)` — the same rule as
|
||
/// [`MaePretrainConfig::validate`]. Erroring (never clamping) keeps the
|
||
/// module's error-not-silent policy: a NaN ratio would otherwise silently
|
||
/// mask zero patches and a ratio ≥ 1 would mask everything.
|
||
pub fn random_mask(n_patches: usize, mask_ratio: f64, seed: u64) -> Result<MaskIndices, MaeError> {
|
||
if !mask_ratio.is_finite() || mask_ratio <= 0.0 || mask_ratio >= 1.0 {
|
||
return Err(MaeError::InvalidMaskRatio { ratio: mask_ratio });
|
||
}
|
||
let n_masked = ((mask_ratio * n_patches as f64).round() as usize).min(n_patches);
|
||
let mut order: Vec<usize> = (0..n_patches).collect();
|
||
let mut rng = Xorshift64::new(seed);
|
||
for i in (1..n_patches).rev() {
|
||
let j = (rng.next_u64() % (i as u64 + 1)) as usize;
|
||
order.swap(i, j);
|
||
}
|
||
let mut masked: Vec<usize> = order[..n_masked].to_vec();
|
||
let mut visible: Vec<usize> = order[n_masked..].to_vec();
|
||
masked.sort_unstable();
|
||
visible.sort_unstable();
|
||
Ok(MaskIndices { masked, visible })
|
||
}
|
||
|
||
// ---------------------------------------------------------------------------
|
||
// helpers
|
||
// ---------------------------------------------------------------------------
|
||
|
||
fn check_axis(axis: &'static str, window: usize, patch: usize) -> Result<(), MaeError> {
|
||
if patch > window {
|
||
return Err(MaeError::PatchExceedsWindow {
|
||
axis,
|
||
patch,
|
||
window,
|
||
});
|
||
}
|
||
let remainder = window % patch;
|
||
if remainder != 0 {
|
||
return Err(MaeError::NotDivisible {
|
||
axis,
|
||
window,
|
||
patch,
|
||
remainder,
|
||
crop: window - remainder,
|
||
});
|
||
}
|
||
Ok(())
|
||
}
|