diff --git a/Cargo.lock b/Cargo.lock index 64883a066..50c34d8c4 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3582,6 +3582,22 @@ dependencies = [ "tokio", ] +[[package]] +name = "pegainfer-higgs-audio" +version = "0.1.0" +dependencies = [ + "anyhow", + "clap", + "half", + "memmap2", + "pegainfer-core", + "pegainfer-qwen3", + "safetensors", + "serde_json", + "sha2 0.11.0", + "tempfile", +] + [[package]] name = "pegainfer-kernels" version = "0.1.0" diff --git a/Cargo.toml b/Cargo.toml index 0baa7bdf6..5f0e1f147 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -15,6 +15,7 @@ members = [ "pegainfer-kimi-k2", "pegainfer-qwen3", "pegainfer-qwen35", + "pegainfer-higgs-audio", "pegainfer-sample", "pegainfer-kv-cache", "pegainfer-kv-offload", @@ -128,6 +129,7 @@ pegainfer-kv-offload = { path = "pegainfer-kv-offload" } pegainfer-kv-store = { path = "pegainfer-kv-store" } pegainfer-qwen3 = { path = "pegainfer-qwen3" } pegainfer-qwen35 = { path = "pegainfer-qwen35" } +pegainfer-higgs-audio = { path = "pegainfer-higgs-audio" } pegainfer-sample = { path = "pegainfer-sample" } opentelemetry = { version = "0.31.0", features = ["trace", "logs"] } opentelemetry-appender-tracing = "0.31.0" diff --git a/docs/models/higgs-audio/a-layer-validation.md b/docs/models/higgs-audio/a-layer-validation.md new file mode 100644 index 000000000..3021b1519 --- /dev/null +++ b/docs/models/higgs-audio/a-layer-validation.md @@ -0,0 +1,92 @@ +# Higgs-Audio A-Layer Validation + +This note records the intended review scope for the Higgs-Audio A-layer bring-up. +It is deliberately narrower than full audio generation: the goal is to load the +Higgs checkpoint/config, run the Qwen3 text backbone prefill, expose traceable +hidden-state boundaries, and verify the one-step audio-head contract against +offline golden tensors. + +## Scope + +Included: + +- Higgs-Audio crate wiring and one-step CLI tools. +- Qwen3 config-view loading backed by tensor-name aliases, without copying a + renamed checkpoint payload. +- Single-GPU Qwen3 diagnostic prefill hooks for final hidden, per-layer hidden, + and selected layer stage dumps. +- One-step audio projection parity tools and a small offline fixture. + +Not included: + +- Full delay-pattern decode state machine. +- Multi-codebook autoregressive audio decode. +- Codec/vocoder integration or waveform output. +- Tensor-parallel diagnostic trace support. + +## Validation Model + +The A-layer validation is golden-trace-driven: + +1. Generate or load a reference prompt and hidden/logit tensors from the Python + or sglang-omni reference stack. +2. Run the PegaInfer Higgs one-step path against the same checkpoint/config. +3. Compare semantic outputs first (`audio_argmax`, cosine similarity, top-k + overlap), then use layer and stage dumps to localize any drift. + +Strict elementwise parity is useful as a diagnostic, but it is not the only +acceptance signal. The current 4090 evidence showed the semantic gate passing +while small BF16/F32 elementwise drift remained: + +The RMSNorm rounding ablation showed that the HF/Qwen3-style fused-add-RMSNorm +variant improves strict trace parity slightly, but is not required for the +current one-step semantic gate. To keep this foundation PR focused, the shared +`pegainfer-kernels` rounding change is left out of scope and can be discussed +separately as a Qwen3 numeric-parity change if needed. + +Evidence source: old-4090 semantic comparison logs from the Higgs-Audio +trace-driven bring-up run; the auto path and retained-session path reported the +same semantic metrics. + +- `audio_argmax.ids`: exact. +- `hidden_cosine`: `0.999994874`. +- `logits_cosine`: `0.999998987`. +- `top64_min_overlap`: `58`. +- `top64_mean_overlap`: `61`. +- `final_hidden.bf16 mean_abs`: about `0.006735`. +- `audio_logits.f32 mean_abs`: about `0.040235`. + +## Local Checks + +Checks that do not require a Linux CUDA runtime: + +```bash +cargo fmt --check -p pegainfer-core -p pegainfer-qwen3 -p pegainfer-higgs-audio +cargo check -p pegainfer-higgs-audio --bins +python3 -m py_compile \ + tools/accuracy/analyze_higgs_one_step_actual.py \ + tools/accuracy/analyze_higgs_projection_drift.py \ + tools/accuracy/analyze_higgs_qk_norm_drift.py \ + tools/accuracy/analyze_higgs_residual_drift.py \ + tools/accuracy/analyze_higgs_rmsnorm_drift.py \ + tools/accuracy/compare_higgs_layer_hidden.py \ + tools/accuracy/compare_higgs_stage_dump.py \ + tools/accuracy/compare_higgs_trace_dump.py \ + tools/accuracy/dump_higgs_layer0_stages_golden.py \ + tools/accuracy/dump_higgs_layer_hidden_golden.py \ + tools/accuracy/dump_higgs_one_step_golden.py \ + tools/higgs/check_higgs_gate_summary.py +``` + +Linux/4090 checks: + +```bash +export PEGAINFER_CUDA_SM=89 +cargo check -p pegainfer-higgs-audio --features runtime-qwen3 --bins +tools/higgs/run_higgs_one_step_cuda_gate.sh +``` + +On macOS, the runtime-Qwen3 check is expected to stop before useful Rust type +checking because the workspace currently builds CUDA kernels and Linux RDMA +dependencies (`rdma-mummy-sys` expects Linux headers such as `endian.h` and +`linux/types.h`). diff --git a/pegainfer-core/src/weight_loader.rs b/pegainfer-core/src/weight_loader.rs index 4b6bc2064..86c8fad03 100644 --- a/pegainfer-core/src/weight_loader.rs +++ b/pegainfer-core/src/weight_loader.rs @@ -24,6 +24,35 @@ use crate::tensor::DeviceContext; use crate::tensor::DeviceMatrix; use crate::tensor::DeviceVec; +/// Optional mapping from the runtime's requested tensor name to the tensor name +/// stored in the safetensors shard. +/// +/// Normal model loading keeps the identity mapping. Model adapters can pass an +/// alias table to reuse a checkpoint layout without materializing a renamed +/// weight copy. +#[derive(Clone, Debug, Default)] +pub struct TensorNameAliases { + storage_by_requested: HashMap, +} + +impl TensorNameAliases { + pub fn new(storage_by_requested: HashMap) -> Self { + Self { + storage_by_requested, + } + } + + pub fn is_empty(&self) -> bool { + self.storage_by_requested.is_empty() + } + + fn storage_name<'a>(&'a self, requested_name: &'a str) -> &'a str { + self.storage_by_requested + .get(requested_name) + .map_or(requested_name, String::as_str) + } +} + mod staging; use staging::ColShardPlan; use staging::WeightStager; @@ -253,18 +282,30 @@ fn find_tensor<'a>( weight_map: &HashMap, name: &str, ) -> Result> { - if let Some(&idx) = weight_map.get(name) { - shards[idx] - .tensor(name) - .map_err(|e| anyhow::anyhow!("Failed to load tensor '{}': {}", name, e)) + find_tensor_with_aliases(shards, weight_map, &TensorNameAliases::default(), name) +} + +fn find_tensor_with_aliases<'a>( + shards: &'a [SafeTensors<'a>], + weight_map: &HashMap, + aliases: &TensorNameAliases, + name: &str, +) -> Result> { + let storage_name = aliases.storage_name(name); + if let Some(&idx) = weight_map.get(storage_name) { + shards[idx].tensor(storage_name).map_err(|e| { + anyhow::anyhow!("Failed to load tensor '{name}' stored as '{storage_name}': {e}") + }) } else { // Fallback: try all shards (single-file case) for shard in shards { - if let Ok(t) = shard.tensor(name) { + if let Ok(t) = shard.tensor(storage_name) { return Ok(t); } } - Err(anyhow::anyhow!("Tensor '{}' not found in any shard", name)) + Err(anyhow::anyhow!( + "Tensor '{name}' stored as '{storage_name}' not found in any shard" + )) } } @@ -361,6 +402,7 @@ pub struct StagedWeightLoader<'a> { stager: WeightStager, shards: &'a [SafeTensors<'a>], weight_map: &'a HashMap, + aliases: TensorNameAliases, slots: Vec, vec_slots: Vec, pending: Vec>, @@ -381,6 +423,7 @@ impl<'a> StagedWeightLoader<'a> { stager: WeightStager::new(ctx)?, shards, weight_map, + aliases: TensorNameAliases::default(), slots: Vec::new(), vec_slots: Vec::new(), pending: Vec::new(), @@ -391,6 +434,12 @@ impl<'a> StagedWeightLoader<'a> { }) } + #[must_use] + pub fn with_aliases(mut self, aliases: TensorNameAliases) -> Self { + self.aliases = aliases; + self + } + fn ensure_recording(&self) -> Result<()> { anyhow::ensure!( !self.finished && !self.failed, @@ -482,7 +531,7 @@ impl<'a> StagedWeightLoader<'a> { } fn tensor_2d(&self, name: &str, rows: usize, cols: usize) -> Result<&'a [u8]> { - let tensor = find_tensor(self.shards, self.weight_map, name)?; + let tensor = find_tensor_with_aliases(self.shards, self.weight_map, &self.aliases, name)?; let shape = tensor.shape(); anyhow::ensure!( shape.len() == 2, @@ -584,7 +633,7 @@ impl<'a> StagedWeightLoader<'a> { /// Small tensors; uploaded as plain pageable copies. pub fn vector(&mut self, name: &str, len: usize) -> Result { self.ensure_recording()?; - let tensor = find_tensor(self.shards, self.weight_map, name)?; + let tensor = find_tensor_with_aliases(self.shards, self.weight_map, &self.aliases, name)?; let shape = tensor.shape(); anyhow::ensure!( shape.len() == 1 && shape[0] == len, diff --git a/pegainfer-higgs-audio/Cargo.toml b/pegainfer-higgs-audio/Cargo.toml new file mode 100644 index 000000000..d55e2f354 --- /dev/null +++ b/pegainfer-higgs-audio/Cargo.toml @@ -0,0 +1,45 @@ +[package] +name = "pegainfer-higgs-audio" +version = "0.1.0" +edition.workspace = true +license.workspace = true + +[dependencies] +anyhow = { workspace = true } +clap = { workspace = true } +half = { workspace = true } +memmap2 = { workspace = true } +pegainfer-core = { workspace = true, optional = true } +pegainfer-qwen3 = { workspace = true, optional = true } +safetensors = { workspace = true } +serde_json = { workspace = true } +sha2 = { workspace = true } + +[features] +runtime-qwen3 = ["dep:pegainfer-core", "dep:pegainfer-qwen3"] + +[dev-dependencies] +tempfile = { workspace = true } + +[[bin]] +name = "higgs_dump_one_step_actual" +path = "src/bin/higgs_dump_one_step_actual.rs" +required-features = ["runtime-qwen3"] + +[[bin]] +name = "higgs_prefill_prompt_session_smoke" +path = "src/bin/higgs_prefill_prompt_session_smoke.rs" +required-features = ["runtime-qwen3"] + +[[bin]] +name = "higgs_dump_prefill_layer_hidden" +path = "src/bin/higgs_dump_prefill_layer_hidden.rs" +required-features = ["runtime-qwen3"] + +[[bin]] +name = "higgs_dump_layer0_stages" +path = "src/bin/higgs_dump_layer0_stages.rs" +required-features = ["runtime-qwen3"] + +[lints] +workspace = true diff --git a/pegainfer-higgs-audio/src/bin/higgs_artifact_check.rs b/pegainfer-higgs-audio/src/bin/higgs_artifact_check.rs new file mode 100644 index 000000000..e74fb2714 --- /dev/null +++ b/pegainfer-higgs-audio/src/bin/higgs_artifact_check.rs @@ -0,0 +1,157 @@ +use std::path::Path; +use std::path::PathBuf; + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use clap::Parser; +use pegainfer_higgs_audio::config::EXPECTED_MODEL_CARD_CONTEXT; +use pegainfer_higgs_audio::config::HiggsConfig; +use pegainfer_higgs_audio::load_plan::HiggsRuntimeLoadPlan; +use pegainfer_higgs_audio::one_step_golden::REQUIRED_TENSORS; +use pegainfer_higgs_audio::one_step_golden::{self}; +use pegainfer_higgs_audio::weights::HiggsWeightManifest; +use pegainfer_higgs_audio::weights::fused_modality_shape; +use pegainfer_higgs_audio::weights::validate_checkpoint_headers; +use sha2::Digest; +use sha2::Sha256; + +#[derive(Parser)] +#[command(about = "Validate the Higgs Audio model artifacts against the one-step golden contract")] +struct Args { + #[arg(long)] + model_dir: PathBuf, + #[arg(long)] + golden: PathBuf, +} + +fn main() -> Result<()> { + let args = Args::parse(); + + let config = HiggsConfig::from_model_dir(&args.model_dir)?; + let manifest = HiggsWeightManifest::from_model_dir(&args.model_dir)?; + let summary = manifest.validate_for_config(&config)?; + let load_plan = HiggsRuntimeLoadPlan::from_manifest(&config, &manifest)?; + let load_summary = load_plan.summary(); + let header_summary = validate_checkpoint_headers(&args.model_dir, &config, &manifest)?; + let golden = one_step_golden::load_and_validate(&args.golden)?; + + check_metadata_hash( + &golden, + "config_sha256", + &args.model_dir.join("config.json"), + )?; + check_metadata_hash( + &golden, + "tokenizer_json_sha256", + &args.model_dir.join("tokenizer.json"), + )?; + check_metadata_hash( + &golden, + "model_index_sha256", + &args.model_dir.join("model.safetensors.index.json"), + )?; + check_optional_model_size(&golden, &args.model_dir.join("model.safetensors"))?; + + println!("higgs artifact check: ok"); + println!(" model_dir: {}", args.model_dir.display()); + println!(" golden: {} bytes sha256={}", golden.bytes, golden.sha256); + println!( + " golden tensors: {} required tensor specs validated", + REQUIRED_TENSORS.len() + ); + println!( + " config: layers={} hidden={} q_heads={} kv_heads={} head_dim={}", + config.text.num_hidden_layers, + config.text.hidden_size, + config.text.num_attention_heads, + config.text.num_key_value_heads, + config.text.head_dim + ); + println!( + " kv bf16: {} bytes/position, {} MiB @ {} positions", + config.kv_bytes_per_position_bf16(), + config.kv_bytes_for_positions_bf16(EXPECTED_MODEL_CARD_CONTEXT) / 1024 / 1024, + EXPECTED_MODEL_CARD_CONTEXT + ); + println!( + " manifest: total={} body={} decoder_only={} fused_modality_shape={:?}", + summary.total_tensors, + summary.body_tensors, + summary.decoder_only_tensors, + fused_modality_shape() + ); + println!( + " checkpoint headers: files={} tensors={} bf16={}", + header_summary.files_checked, + header_summary.tensors_checked, + header_summary.bf16_tensors_checked + ); + println!( + " runtime load plan: tensors={} shard_files={} bf16_mib={} qwen3_backbone={} higgs_head={}", + load_summary.tensors, + load_summary.shard_files, + load_summary.bf16_bytes / 1024 / 1024, + load_summary.qwen3_backbone_tensors, + load_summary.higgs_head_tensors + ); + + Ok(()) +} + +fn check_metadata_hash( + golden: &one_step_golden::GoldenContract, + metadata_key: &str, + path: &Path, +) -> Result<()> { + let expected = golden + .metadata + .get(metadata_key) + .with_context(|| format!("golden metadata missing {metadata_key}"))?; + let actual = sha256_file(path)?; + if &actual != expected { + bail!( + "{metadata_key} mismatch for {}: expected {expected}, got {actual}", + path.display() + ); + } + Ok(()) +} + +fn check_optional_model_size( + golden: &one_step_golden::GoldenContract, + model_safetensors: &Path, +) -> Result<()> { + let Some(expected) = golden.metadata.get("model_safetensors_size") else { + return Ok(()); + }; + if !model_safetensors.exists() { + println!( + " note: skipping model.safetensors size check because {} is absent", + model_safetensors.display() + ); + return Ok(()); + } + let actual = std::fs::metadata(model_safetensors) + .with_context(|| format!("stat {}", model_safetensors.display()))? + .len() + .to_string(); + if &actual != expected { + bail!( + "model_safetensors_size mismatch for {}: expected {expected}, got {actual}", + model_safetensors.display() + ); + } + Ok(()) +} + +fn sha256_file(path: &Path) -> Result { + let bytes = std::fs::read(path).with_context(|| format!("read {}", path.display()))?; + let mut digest = Sha256::new(); + digest.update(bytes); + Ok(digest + .finalize() + .iter() + .map(|byte| format!("{byte:02x}")) + .collect()) +} diff --git a/pegainfer-higgs-audio/src/bin/higgs_compare_one_step.rs b/pegainfer-higgs-audio/src/bin/higgs_compare_one_step.rs new file mode 100644 index 000000000..4799cb66c --- /dev/null +++ b/pegainfer-higgs-audio/src/bin/higgs_compare_one_step.rs @@ -0,0 +1,124 @@ +use std::path::PathBuf; + +use anyhow::Result; +use clap::Parser; +use clap::ValueEnum; +use pegainfer_higgs_audio::compare::OneStepSemanticTolerances; +use pegainfer_higgs_audio::compare::OneStepTolerances; +use pegainfer_higgs_audio::compare::compare_one_step_files; +use pegainfer_higgs_audio::compare::compare_one_step_semantic_files; +use pegainfer_higgs_audio::compare::ensure_comparison_passed; +use pegainfer_higgs_audio::compare::ensure_semantic_comparison_passed; + +#[derive(Clone, Copy, Debug, Eq, PartialEq, ValueEnum)] +enum CompareMode { + /// Enforce exact tensor identity plus tight numeric tolerances. + Strict, + /// Print strict drift diagnostics, then enforce semantic runtime parity. + Semantic, +} + +#[derive(Parser)] +#[command(about = "Compare a Higgs Audio one-step actual safetensors dump against the golden")] +struct Args { + #[arg(long)] + golden: PathBuf, + #[arg(long)] + actual: PathBuf, + #[arg(long, value_enum, default_value_t = CompareMode::Strict)] + mode: CompareMode, + #[arg(long, default_value_t = OneStepTolerances::default().hidden_abs_tol)] + hidden_abs_tol: f32, + #[arg(long, default_value_t = OneStepTolerances::default().hidden_mean_abs_tol)] + hidden_mean_abs_tol: f32, + #[arg(long, default_value_t = OneStepTolerances::default().logits_abs_tol)] + logits_abs_tol: f32, + #[arg(long, default_value_t = OneStepTolerances::default().logits_mean_abs_tol)] + logits_mean_abs_tol: f32, + #[arg(long, default_value_t = OneStepTolerances::default().top_logprobs_abs_tol)] + top_logprobs_abs_tol: f32, + #[arg(long, default_value_t = OneStepTolerances::default().top_logprobs_mean_abs_tol)] + top_logprobs_mean_abs_tol: f32, + #[arg(long, default_value_t = OneStepSemanticTolerances::default().hidden_cosine_min)] + hidden_cosine_min: f32, + #[arg(long, default_value_t = OneStepSemanticTolerances::default().logits_cosine_min)] + logits_cosine_min: f32, + #[arg(long, default_value_t = OneStepSemanticTolerances::default().argmax_regret_tol)] + argmax_regret_tol: f32, + #[arg(long, default_value_t = OneStepSemanticTolerances::default().top64_min_overlap)] + top64_min_overlap: usize, +} + +fn main() -> Result<()> { + let args = Args::parse(); + let tolerances = OneStepTolerances { + hidden_abs_tol: args.hidden_abs_tol, + hidden_mean_abs_tol: args.hidden_mean_abs_tol, + logits_abs_tol: args.logits_abs_tol, + logits_mean_abs_tol: args.logits_mean_abs_tol, + top_logprobs_abs_tol: args.top_logprobs_abs_tol, + top_logprobs_mean_abs_tol: args.top_logprobs_mean_abs_tol, + }; + let comparison = compare_one_step_files(&args.golden, &args.actual, tolerances)?; + println!("higgs one-step strict comparison:"); + for tensor in &comparison.tensors { + println!( + " {:32} pass={} elems={} exact_mismatch={} max_abs={:.6} mean_abs={:.6} rmse={:.6} p99_abs={:.6} abs_tol={:.6} mean_tol={:.6}", + tensor.name, + tensor.passed, + tensor.elements, + tensor.exact_mismatches, + tensor.max_abs, + tensor.mean_abs, + tensor.rmse, + tensor.p99_abs, + tensor.abs_tol, + tensor.mean_abs_tol + ); + } + + match args.mode { + CompareMode::Strict => { + ensure_comparison_passed(&comparison)?; + println!("higgs one-step strict comparison: ok"); + } + CompareMode::Semantic => { + println!( + "higgs one-step strict comparison: passed={} diagnostic_only=true", + comparison.passed() + ); + let semantic_tolerances = OneStepSemanticTolerances { + hidden_cosine_min: args.hidden_cosine_min, + logits_cosine_min: args.logits_cosine_min, + argmax_regret_tol: args.argmax_regret_tol, + top64_min_overlap: args.top64_min_overlap, + }; + let semantic = + compare_one_step_semantic_files(&args.golden, &args.actual, semantic_tolerances)?; + println!("higgs one-step semantic comparison:"); + println!( + " prompt_exact={} argmax_exact={} hidden_cosine={:.9} hidden_cosine_min={:.9}", + semantic.prompt_exact, + semantic.audio_argmax_exact, + semantic.hidden_cosine, + semantic.tolerances.hidden_cosine_min + ); + println!( + " logits_cosine={:.9} logits_cosine_min={:.9} max_argmax_regret={:.6} argmax_regret_tol={:.6}", + semantic.logits_cosine, + semantic.tolerances.logits_cosine_min, + semantic.max_argmax_regret, + semantic.tolerances.argmax_regret_tol + ); + println!( + " top64_min_overlap={} top64_mean_overlap={:.2} top64_min_overlap_tol={}", + semantic.top64_min_overlap, + semantic.top64_mean_overlap, + semantic.tolerances.top64_min_overlap + ); + ensure_semantic_comparison_passed(&semantic)?; + println!("higgs one-step semantic comparison: ok"); + } + } + Ok(()) +} diff --git a/pegainfer-higgs-audio/src/bin/higgs_dump_layer0_stages.rs b/pegainfer-higgs-audio/src/bin/higgs_dump_layer0_stages.rs new file mode 100644 index 000000000..8f3822e57 --- /dev/null +++ b/pegainfer-higgs-audio/src/bin/higgs_dump_layer0_stages.rs @@ -0,0 +1,49 @@ +use std::path::PathBuf; + +use anyhow::Result; +use clap::Parser; +use pegainfer_higgs_audio::layer_dump::write_stage_dump; +use pegainfer_higgs_audio::one_step_actual::load_prompt_from_golden; +use pegainfer_qwen3::runtime::Qwen3Executor; + +#[derive(Parser)] +#[command( + about = "Dump Higgs/Qwen3 prefill stage snapshots for one selected layer and one golden prompt" +)] +struct Args { + /// Qwen3-compatible body view produced by higgs_materialize_qwen3_body. + #[arg(long)] + qwen3_body_dir: PathBuf, + /// Golden safetensors fixture; prompt tensors are copied from this file. + #[arg(long)] + golden: PathBuf, + /// Output layer stage safetensors path. + #[arg(long)] + out: PathBuf, + /// Zero-based decoder layer index to stage-dump. + #[arg(long, default_value_t = 0)] + layer_idx: usize, + #[arg(long, default_value_t = 0)] + device_ordinal: usize, +} + +fn main() -> Result<()> { + let args = Args::parse(); + let prompt = load_prompt_from_golden(&args.golden)?; + let prompt_ids = prompt.prompt_ids()?; + let qwen3_body_dir = args + .qwen3_body_dir + .to_str() + .ok_or_else(|| anyhow::anyhow!("qwen3 body dir must be valid UTF-8"))?; + let mut executor = Qwen3Executor::from_runtime(qwen3_body_dir, false, &[args.device_ordinal])?; + let stages = executor.prefill_layer_stages_bf16(args.layer_idx, prompt_ids)?; + let summary = write_stage_dump(&args.out, &prompt, &stages.stages)?; + + println!("higgs layer stage dump: ok"); + println!(" out: {}", summary.output_path.display()); + println!(" layer_idx: {}", args.layer_idx); + println!(" prompt_tokens: {}", summary.prompt_tokens); + println!(" stages: {}", summary.stages); + println!(" values: {}", summary.values); + Ok(()) +} diff --git a/pegainfer-higgs-audio/src/bin/higgs_dump_one_step_actual.rs b/pegainfer-higgs-audio/src/bin/higgs_dump_one_step_actual.rs new file mode 100644 index 000000000..7d4a5e6ba --- /dev/null +++ b/pegainfer-higgs-audio/src/bin/higgs_dump_one_step_actual.rs @@ -0,0 +1,87 @@ +use std::path::PathBuf; + +use anyhow::Result; +use clap::Parser; +use clap::ValueEnum; +use pegainfer_higgs_audio::runtime_bridge::AudioHeadBackend as RuntimeAudioHeadBackend; +use pegainfer_higgs_audio::runtime_bridge::HiggsAudioRuntime; +use pegainfer_higgs_audio::runtime_bridge::HiggsRuntimeSource; +use pegainfer_higgs_audio::runtime_source::Qwen3RuntimeSourcePath; +use pegainfer_higgs_audio::runtime_source::select_qwen3_runtime_source; + +#[derive(Clone, Copy, Debug, Eq, PartialEq, ValueEnum)] +enum AudioHeadBackend { + /// Match the Python golden generator's CUDA BF16 F.linear contract. + CudaBf16, + /// Legacy diagnostic fallback: compute the audio head as a CPU FP32 dot. + CpuFp32, +} + +#[derive(Parser)] +#[command(about = "Dump a Higgs Audio one-step actual safetensors file from PegaInfer runtime")] +struct Args { + /// Original Higgs checkpoint directory containing the fused audio head. + #[arg(long)] + model_dir: PathBuf, + /// Optional fallback Qwen3-compatible body view produced by higgs_materialize_qwen3_body. + #[arg(long, conflicts_with = "qwen3_config_dir")] + qwen3_body_dir: Option, + /// Optional Qwen3 config-only view; by default a small view is written next to --out. + #[arg(long, conflicts_with = "qwen3_body_dir")] + qwen3_config_dir: Option, + /// Golden safetensors fixture; prompt tensors are copied from this file. + #[arg(long)] + golden: PathBuf, + /// Output actual safetensors path. + #[arg(long)] + out: PathBuf, + #[arg(long, default_value_t = 0)] + device_ordinal: usize, + /// Audio head execution backend used for the actual logits dump. + #[arg(long, value_enum, default_value_t = AudioHeadBackend::CudaBf16)] + audio_head_backend: AudioHeadBackend, +} + +fn main() -> Result<()> { + let args = Args::parse(); + let source_path = select_qwen3_runtime_source( + args.qwen3_body_dir.as_deref(), + args.qwen3_config_dir.as_deref(), + &args.out, + )?; + let source = match &source_path { + Qwen3RuntimeSourcePath::BodyView(qwen3_body_dir) => { + HiggsRuntimeSource::Qwen3BodyView { qwen3_body_dir } + } + Qwen3RuntimeSourcePath::ConfigAlias(qwen3_config_dir) => { + HiggsRuntimeSource::Qwen3ConfigAlias { qwen3_config_dir } + } + Qwen3RuntimeSourcePath::AutoConfigAlias(qwen3_config_dir) => { + HiggsRuntimeSource::AutoConfigAlias { qwen3_config_dir } + } + }; + let mut runtime = HiggsAudioRuntime::from_model_dir( + &args.model_dir, + source, + args.audio_head_backend.into(), + args.device_ordinal, + )?; + let summary = runtime.dump_one_step_actual(&args.golden, &args.out)?; + + println!("higgs one-step actual dump: ok"); + println!(" out: {}", summary.output_path.display()); + println!(" audio_head_backend: {:?}", args.audio_head_backend); + println!(" prompt_tokens: {}", summary.prompt_tokens); + println!(" hidden_values: {}", summary.hidden_values); + println!(" audio_logits: {}", summary.audio_logits); + Ok(()) +} + +impl From for RuntimeAudioHeadBackend { + fn from(value: AudioHeadBackend) -> Self { + match value { + AudioHeadBackend::CudaBf16 => Self::CudaBf16, + AudioHeadBackend::CpuFp32 => Self::CpuFp32, + } + } +} diff --git a/pegainfer-higgs-audio/src/bin/higgs_dump_prefill_layer_hidden.rs b/pegainfer-higgs-audio/src/bin/higgs_dump_prefill_layer_hidden.rs new file mode 100644 index 000000000..3eb473e44 --- /dev/null +++ b/pegainfer-higgs-audio/src/bin/higgs_dump_prefill_layer_hidden.rs @@ -0,0 +1,52 @@ +use std::path::PathBuf; + +use anyhow::Result; +use clap::Parser; +use pegainfer_higgs_audio::layer_dump::write_layer_hidden_dump; +use pegainfer_higgs_audio::one_step_actual::load_prompt_from_golden; +use pegainfer_qwen3::runtime::Qwen3Executor; + +#[derive(Parser)] +#[command(about = "Dump Higgs/Qwen3 per-layer prefill hidden snapshots for one golden prompt")] +struct Args { + /// Qwen3-compatible body view produced by higgs_materialize_qwen3_body. + #[arg(long)] + qwen3_body_dir: PathBuf, + /// Golden safetensors fixture; prompt tensors are copied from this file. + #[arg(long)] + golden: PathBuf, + /// Output layer hidden safetensors path. + #[arg(long)] + out: PathBuf, + #[arg(long, default_value_t = 0)] + device_ordinal: usize, +} + +fn main() -> Result<()> { + let args = Args::parse(); + let prompt = load_prompt_from_golden(&args.golden)?; + let prompt_ids = prompt.prompt_ids()?; + let qwen3_body_dir = args + .qwen3_body_dir + .to_str() + .ok_or_else(|| anyhow::anyhow!("qwen3 body dir must be valid UTF-8"))?; + let mut executor = Qwen3Executor::from_runtime(qwen3_body_dir, false, &[args.device_ordinal])?; + let hidden = executor.prefill_layer_hidden_bf16(prompt_ids)?; + let summary = write_layer_hidden_dump( + &args.out, + &prompt, + &hidden.embedding_hidden_bf16, + &hidden.layer_hidden_bf16, + &hidden.final_normed_bf16, + )?; + + println!("higgs prefill layer hidden dump: ok"); + println!(" out: {}", summary.output_path.display()); + println!(" prompt_tokens: {}", summary.prompt_tokens); + println!(" layers: {}", summary.layers); + println!( + " hidden_values_per_layer: {}", + summary.hidden_values_per_layer + ); + Ok(()) +} diff --git a/pegainfer-higgs-audio/src/bin/higgs_materialize_qwen3_body.rs b/pegainfer-higgs-audio/src/bin/higgs_materialize_qwen3_body.rs new file mode 100644 index 000000000..2054cc870 --- /dev/null +++ b/pegainfer-higgs-audio/src/bin/higgs_materialize_qwen3_body.rs @@ -0,0 +1,45 @@ +use std::path::PathBuf; + +use anyhow::Result; +use clap::Parser; +use pegainfer_higgs_audio::config::HiggsConfig; +use pegainfer_higgs_audio::load_plan::HiggsRuntimeLoadPlan; +use pegainfer_higgs_audio::materialize_qwen3::materialize_qwen3_body_view; +use pegainfer_higgs_audio::materialize_qwen3::write_qwen3_config_view; +use pegainfer_higgs_audio::weights::HiggsWeightManifest; +use pegainfer_higgs_audio::weights::validate_checkpoint_headers; + +#[derive(Parser)] +#[command(about = "Materialize a Qwen3-compatible view of the Higgs text/body checkpoint")] +struct Args { + #[arg(long)] + model_dir: PathBuf, + #[arg(long)] + out_dir: PathBuf, + /// Write only Qwen3 config files plus a tensor-alias manifest; do not copy weight payloads. + #[arg(long)] + metadata_only: bool, +} + +fn main() -> Result<()> { + let args = Args::parse(); + let config = HiggsConfig::from_model_dir(&args.model_dir)?; + let manifest = HiggsWeightManifest::from_model_dir(&args.model_dir)?; + let plan = HiggsRuntimeLoadPlan::from_manifest(&config, &manifest)?; + validate_checkpoint_headers(&args.model_dir, &config, &manifest)?; + if args.metadata_only { + let summary = write_qwen3_config_view(&args.out_dir, &config, &plan)?; + println!("higgs qwen3 config view materialized: ok"); + println!(" out_dir: {}", summary.output_dir.display()); + println!(" alias_manifest: {}", summary.alias_manifest.display()); + println!(" aliases: {}", summary.aliases); + return Ok(()); + } + + let summary = materialize_qwen3_body_view(&args.model_dir, &args.out_dir, &config, &plan)?; + println!("higgs qwen3 body view materialized: ok"); + println!(" out_dir: {}", summary.output_dir.display()); + println!(" tensors: {}", summary.tensors); + println!(" payload_mib: {}", summary.payload_bytes / 1024 / 1024); + Ok(()) +} diff --git a/pegainfer-higgs-audio/src/bin/higgs_prefill_prompt_session_smoke.rs b/pegainfer-higgs-audio/src/bin/higgs_prefill_prompt_session_smoke.rs new file mode 100644 index 000000000..8783ea91d --- /dev/null +++ b/pegainfer-higgs-audio/src/bin/higgs_prefill_prompt_session_smoke.rs @@ -0,0 +1,113 @@ +use std::path::PathBuf; + +use anyhow::Result; +use anyhow::bail; +use clap::Parser; +use clap::ValueEnum; +use pegainfer_higgs_audio::one_step_actual::load_prompt_from_golden; +use pegainfer_higgs_audio::one_step_actual::write_one_step_actual_prediction; +use pegainfer_higgs_audio::runtime_bridge::AudioHeadBackend as RuntimeAudioHeadBackend; +use pegainfer_higgs_audio::runtime_bridge::HiggsAudioRuntime; +use pegainfer_higgs_audio::runtime_bridge::HiggsPromptSession; +use pegainfer_higgs_audio::runtime_bridge::HiggsRuntimeSource; +use pegainfer_higgs_audio::runtime_source::Qwen3RuntimeSourcePath; +use pegainfer_higgs_audio::runtime_source::select_qwen3_runtime_source; + +#[derive(Clone, Copy, Debug, Eq, PartialEq, ValueEnum)] +enum AudioHeadBackend { + CudaBf16, + CpuFp32, +} + +#[derive(Parser)] +#[command(about = "Smoke-test Higgs prompt-session prefill and write a one-step actual file")] +struct Args { + /// Original Higgs checkpoint directory containing the fused audio head. + #[arg(long)] + model_dir: PathBuf, + /// Optional fallback Qwen3-compatible body view produced by higgs_materialize_qwen3_body. + #[arg(long, conflicts_with = "qwen3_config_dir")] + qwen3_body_dir: Option, + /// Optional Qwen3 config-only view; by default a small view is written next to --out. + #[arg(long, conflicts_with = "qwen3_body_dir")] + qwen3_config_dir: Option, + /// Golden safetensors fixture; prompt tensors are copied from this file. + #[arg(long)] + golden: PathBuf, + /// Output actual safetensors path generated from the retained prompt session. + #[arg(long)] + out: PathBuf, + #[arg(long, default_value_t = 1)] + request_id: u64, + #[arg(long, default_value_t = 0)] + device_ordinal: usize, + /// Audio head execution backend used for the actual logits dump. + #[arg(long, value_enum, default_value_t = AudioHeadBackend::CudaBf16)] + audio_head_backend: AudioHeadBackend, +} + +fn main() -> Result<()> { + let args = Args::parse(); + let source_path = select_qwen3_runtime_source( + args.qwen3_body_dir.as_deref(), + args.qwen3_config_dir.as_deref(), + &args.out, + )?; + let source = match &source_path { + Qwen3RuntimeSourcePath::BodyView(qwen3_body_dir) => { + HiggsRuntimeSource::Qwen3BodyView { qwen3_body_dir } + } + Qwen3RuntimeSourcePath::ConfigAlias(qwen3_config_dir) => { + HiggsRuntimeSource::Qwen3ConfigAlias { qwen3_config_dir } + } + Qwen3RuntimeSourcePath::AutoConfigAlias(qwen3_config_dir) => { + HiggsRuntimeSource::AutoConfigAlias { qwen3_config_dir } + } + }; + + let prompt = load_prompt_from_golden(&args.golden)?; + let prompt_ids = prompt.prompt_ids()?; + let mut runtime = HiggsAudioRuntime::from_model_dir( + &args.model_dir, + source, + args.audio_head_backend.into(), + args.device_ordinal, + )?; + let session_handle = HiggsPromptSession::new(args.request_id); + let session = runtime.prefill_prompt_session(session_handle, &prompt_ids)?; + if runtime + .prefill_prompt_session(session_handle, &prompt_ids) + .is_ok() + { + bail!( + "duplicate Higgs prompt-session prefill unexpectedly replaced request_id={}", + session_handle.id() + ); + } + let summary = write_one_step_actual_prediction( + &args.out, + &prompt, + &session.final_hidden_bf16, + &session.audio, + )?; + runtime.drop_prompt_session(session_handle)?; + + println!("higgs prompt-session prefill smoke: ok"); + println!(" request_id: {}", session.session.id()); + println!(" duplicate_request_id_guard: ok"); + println!(" out: {}", summary.output_path.display()); + println!(" audio_head_backend: {:?}", args.audio_head_backend); + println!(" prompt_tokens: {}", summary.prompt_tokens); + println!(" hidden_values: {}", summary.hidden_values); + println!(" audio_logits: {}", summary.audio_logits); + Ok(()) +} + +impl From for RuntimeAudioHeadBackend { + fn from(value: AudioHeadBackend) -> Self { + match value { + AudioHeadBackend::CudaBf16 => Self::CudaBf16, + AudioHeadBackend::CpuFp32 => Self::CpuFp32, + } + } +} diff --git a/pegainfer-higgs-audio/src/compare.rs b/pegainfer-higgs-audio/src/compare.rs new file mode 100644 index 000000000..b14d47f0e --- /dev/null +++ b/pegainfer-higgs-audio/src/compare.rs @@ -0,0 +1,1054 @@ +use std::path::Path; + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use anyhow::ensure; +use half::bf16; +use safetensors::Dtype; +use safetensors::SafeTensors; +use safetensors::tensor::TensorView; + +use crate::one_step_golden::CODEBOOK_VOCAB_SIZE; +use crate::one_step_golden::HIDDEN_SIZE; +use crate::one_step_golden::NUM_CODEBOOKS; +use crate::one_step_golden::TOP_K; + +pub const PROMPT_INPUT_IDS: &str = "prompt.input_ids_padded"; +pub const PROMPT_ATTENTION_MASK: &str = "prompt.attention_mask"; +pub const PROMPT_LENGTHS: &str = "prompt.lengths"; +pub const FINAL_HIDDEN_BF16: &str = "final_hidden.bf16"; +pub const AUDIO_LOGITS_F32: &str = "audio_logits.f32"; +pub const AUDIO_TOP64_IDS: &str = "audio_top64.ids"; +pub const AUDIO_TOP64_LOGPROBS_F32: &str = "audio_top64.logprobs.f32"; +pub const AUDIO_ARGMAX_IDS: &str = "audio_argmax.ids"; + +#[derive(Debug, Clone, Copy, PartialEq)] +pub struct OneStepTolerances { + pub hidden_abs_tol: f32, + pub hidden_mean_abs_tol: f32, + pub logits_abs_tol: f32, + pub logits_mean_abs_tol: f32, + pub top_logprobs_abs_tol: f32, + pub top_logprobs_mean_abs_tol: f32, +} + +impl Default for OneStepTolerances { + fn default() -> Self { + Self { + hidden_abs_tol: 0.03125, + hidden_mean_abs_tol: 0.003, + logits_abs_tol: 0.05, + logits_mean_abs_tol: 0.005, + top_logprobs_abs_tol: 0.05, + top_logprobs_mean_abs_tol: 0.005, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq)] +pub struct OneStepSemanticTolerances { + pub hidden_cosine_min: f32, + pub logits_cosine_min: f32, + pub argmax_regret_tol: f32, + pub top64_min_overlap: usize, +} + +impl Default for OneStepSemanticTolerances { + fn default() -> Self { + Self { + hidden_cosine_min: 0.9998, + logits_cosine_min: 0.99999, + argmax_regret_tol: 0.20, + top64_min_overlap: 40, + } + } +} + +#[derive(Debug, Clone, PartialEq)] +pub struct TensorComparison { + pub name: &'static str, + pub elements: usize, + pub exact_mismatches: usize, + pub max_abs: f32, + pub mean_abs: f32, + pub rmse: f32, + pub p99_abs: f32, + pub abs_tol: f32, + pub mean_abs_tol: f32, + pub passed: bool, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct OneStepComparison { + pub tensors: Vec, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct OneStepSemanticComparison { + pub prompt_exact: bool, + pub audio_argmax_exact: bool, + pub hidden_cosine: f32, + pub logits_cosine: f32, + pub max_argmax_regret: f32, + pub top64_min_overlap: usize, + pub top64_mean_overlap: f32, + pub tolerances: OneStepSemanticTolerances, +} + +impl OneStepComparison { + pub fn passed(&self) -> bool { + self.tensors.iter().all(|tensor| tensor.passed) + } +} + +impl OneStepSemanticComparison { + pub fn passed(&self) -> bool { + self.prompt_exact + && self.audio_argmax_exact + && self.hidden_cosine >= self.tolerances.hidden_cosine_min + && self.logits_cosine >= self.tolerances.logits_cosine_min + && self.max_argmax_regret <= self.tolerances.argmax_regret_tol + && self.top64_min_overlap >= self.tolerances.top64_min_overlap + } +} + +pub fn compare_one_step_files( + golden: impl AsRef, + actual: impl AsRef, + tolerances: OneStepTolerances, +) -> Result { + let golden_path = golden.as_ref(); + let actual_path = actual.as_ref(); + let golden_bytes = + std::fs::read(golden_path).with_context(|| format!("read {}", golden_path.display()))?; + let actual_bytes = + std::fs::read(actual_path).with_context(|| format!("read {}", actual_path.display()))?; + let golden_st = + SafeTensors::deserialize(&golden_bytes).context("parse Higgs golden safetensors")?; + let actual_st = + SafeTensors::deserialize(&actual_bytes).context("parse Higgs actual safetensors")?; + compare_one_step_safetensors(&golden_st, &actual_st, tolerances) +} + +pub fn compare_one_step_semantic_files( + golden: impl AsRef, + actual: impl AsRef, + tolerances: OneStepSemanticTolerances, +) -> Result { + let golden_path = golden.as_ref(); + let actual_path = actual.as_ref(); + let golden_bytes = + std::fs::read(golden_path).with_context(|| format!("read {}", golden_path.display()))?; + let actual_bytes = + std::fs::read(actual_path).with_context(|| format!("read {}", actual_path.display()))?; + let golden_st = + SafeTensors::deserialize(&golden_bytes).context("parse Higgs golden safetensors")?; + let actual_st = + SafeTensors::deserialize(&actual_bytes).context("parse Higgs actual safetensors")?; + compare_one_step_semantic_safetensors(&golden_st, &actual_st, tolerances) +} + +pub fn compare_one_step_safetensors( + golden: &SafeTensors, + actual: &SafeTensors, + tolerances: OneStepTolerances, +) -> Result { + validate_comparison_schema(golden, actual)?; + + let tensors = vec![ + compare_i64_exact(golden, actual, PROMPT_INPUT_IDS)?, + compare_i64_exact(golden, actual, PROMPT_ATTENTION_MASK)?, + compare_i64_exact(golden, actual, PROMPT_LENGTHS)?, + compare_bf16( + golden, + actual, + FINAL_HIDDEN_BF16, + tolerances.hidden_abs_tol, + tolerances.hidden_mean_abs_tol, + )?, + compare_f32( + golden, + actual, + AUDIO_LOGITS_F32, + tolerances.logits_abs_tol, + tolerances.logits_mean_abs_tol, + )?, + compare_i64_exact(golden, actual, AUDIO_TOP64_IDS)?, + compare_f32( + golden, + actual, + AUDIO_TOP64_LOGPROBS_F32, + tolerances.top_logprobs_abs_tol, + tolerances.top_logprobs_mean_abs_tol, + )?, + compare_i64_exact(golden, actual, AUDIO_ARGMAX_IDS)?, + ]; + + Ok(OneStepComparison { tensors }) +} + +pub fn compare_one_step_semantic_safetensors( + golden: &SafeTensors, + actual: &SafeTensors, + tolerances: OneStepSemanticTolerances, +) -> Result { + validate_comparison_schema(golden, actual)?; + + let prompt_exact = i64_values(tensor(golden, PROMPT_INPUT_IDS)?)? + == i64_values(tensor(actual, PROMPT_INPUT_IDS)?)? + && i64_values(tensor(golden, PROMPT_ATTENTION_MASK)?)? + == i64_values(tensor(actual, PROMPT_ATTENTION_MASK)?)? + && i64_values(tensor(golden, PROMPT_LENGTHS)?)? + == i64_values(tensor(actual, PROMPT_LENGTHS)?)?; + let golden_argmax = i64_values(tensor(golden, AUDIO_ARGMAX_IDS)?)?; + let actual_argmax = i64_values(tensor(actual, AUDIO_ARGMAX_IDS)?)?; + let audio_argmax_exact = golden_argmax == actual_argmax; + + let golden_hidden = bf16_values(tensor(golden, FINAL_HIDDEN_BF16)?)?; + let actual_hidden = bf16_values(tensor(actual, FINAL_HIDDEN_BF16)?)?; + let hidden_cosine = cosine_similarity(&golden_hidden, &actual_hidden)?; + + let golden_logits = f32_values(tensor(golden, AUDIO_LOGITS_F32)?)?; + let actual_logits = f32_values(tensor(actual, AUDIO_LOGITS_F32)?)?; + let logits_cosine = cosine_similarity(&golden_logits, &actual_logits)?; + + let max_argmax_regret = max_argmax_regret(&golden_logits, &actual_argmax)?; + let (top64_min_overlap, top64_mean_overlap) = + topk_overlap_from_logits(&golden_logits, &actual_logits, actual_argmax.len(), 64)?; + + Ok(OneStepSemanticComparison { + prompt_exact, + audio_argmax_exact, + hidden_cosine, + logits_cosine, + max_argmax_regret, + top64_min_overlap, + top64_mean_overlap, + tolerances, + }) +} + +pub fn ensure_comparison_passed(comparison: &OneStepComparison) -> Result<()> { + if comparison.passed() { + return Ok(()); + } + let failing: Vec<_> = comparison + .tensors + .iter() + .filter(|tensor| !tensor.passed) + .map(|tensor| tensor.name) + .collect(); + bail!("Higgs one-step comparison failed for tensor(s): {failing:?}"); +} + +pub fn ensure_semantic_comparison_passed(comparison: &OneStepSemanticComparison) -> Result<()> { + if comparison.passed() { + return Ok(()); + } + bail!( + "Higgs one-step semantic comparison failed: prompt_exact={} argmax_exact={} hidden_cosine={:.9} logits_cosine={:.9} max_argmax_regret={:.6} top64_min_overlap={} top64_mean_overlap={:.2}", + comparison.prompt_exact, + comparison.audio_argmax_exact, + comparison.hidden_cosine, + comparison.logits_cosine, + comparison.max_argmax_regret, + comparison.top64_min_overlap, + comparison.top64_mean_overlap + ); +} + +fn compare_i64_exact( + golden: &SafeTensors, + actual: &SafeTensors, + name: &'static str, +) -> Result { + let golden = tensor(golden, name)?; + let actual = tensor(actual, name)?; + ensure!(golden.dtype() == Dtype::I64, "{name} golden must be I64"); + ensure!(actual.dtype() == Dtype::I64, "{name} actual must be I64"); + let golden_values = i64_values(golden)?; + let actual_values = i64_values(actual)?; + ensure!( + golden_values.len() == actual_values.len(), + "{name} element count mismatch: golden {} actual {}", + golden_values.len(), + actual_values.len() + ); + + let diffs: Vec = golden_values + .iter() + .zip(&actual_values) + .map(|(golden, actual)| (*golden - *actual).unsigned_abs() as f32) + .collect(); + let exact_mismatches = diffs.iter().filter(|diff| **diff != 0.0).count(); + let stats = stats(&diffs); + Ok(TensorComparison { + name, + elements: diffs.len(), + exact_mismatches, + max_abs: stats.max_abs, + mean_abs: stats.mean_abs, + rmse: stats.rmse, + p99_abs: stats.p99_abs, + abs_tol: 0.0, + mean_abs_tol: 0.0, + passed: exact_mismatches == 0, + }) +} + +fn compare_f32( + golden: &SafeTensors, + actual: &SafeTensors, + name: &'static str, + abs_tol: f32, + mean_abs_tol: f32, +) -> Result { + let golden = tensor(golden, name)?; + let actual = tensor(actual, name)?; + ensure!(golden.dtype() == Dtype::F32, "{name} golden must be F32"); + ensure!(actual.dtype() == Dtype::F32, "{name} actual must be F32"); + compare_float_values( + name, + &f32_values(golden)?, + &f32_values(actual)?, + abs_tol, + mean_abs_tol, + ) +} + +fn validate_comparison_schema(golden: &SafeTensors, actual: &SafeTensors) -> Result<()> { + let prompt_shape = require_matching_tensor(golden, actual, PROMPT_INPUT_IDS, Dtype::I64)?; + ensure!( + prompt_shape.len() == 2, + "{PROMPT_INPUT_IDS} must be rank-2 [batch, seq], got {prompt_shape:?}" + ); + let batch = prompt_shape[0]; + let seq = prompt_shape[1]; + ensure!(batch > 0, "{PROMPT_INPUT_IDS} batch must be non-zero"); + ensure!(seq > 0, "{PROMPT_INPUT_IDS} seq must be non-zero"); + + let attention_shape = + require_matching_tensor(golden, actual, PROMPT_ATTENTION_MASK, Dtype::I64)?; + ensure!( + attention_shape == prompt_shape, + "{PROMPT_ATTENTION_MASK} shape {attention_shape:?} must match {PROMPT_INPUT_IDS} shape {prompt_shape:?}" + ); + + let lengths_shape = require_matching_tensor(golden, actual, PROMPT_LENGTHS, Dtype::I64)?; + ensure!( + lengths_shape == [batch], + "{PROMPT_LENGTHS} shape {lengths_shape:?} must be [batch={batch}]" + ); + validate_prompt_surface(golden, "golden", batch, seq)?; + validate_prompt_surface(actual, "actual", batch, seq)?; + + let hidden_shape = require_matching_tensor(golden, actual, FINAL_HIDDEN_BF16, Dtype::BF16)?; + ensure!( + hidden_shape == [batch, HIDDEN_SIZE], + "{FINAL_HIDDEN_BF16} shape {hidden_shape:?} must be [batch={batch}, hidden={HIDDEN_SIZE}]" + ); + + let logits_shape = require_matching_tensor(golden, actual, AUDIO_LOGITS_F32, Dtype::F32)?; + ensure!( + logits_shape == [batch, NUM_CODEBOOKS, CODEBOOK_VOCAB_SIZE], + "{AUDIO_LOGITS_F32} shape {logits_shape:?} must be [batch={batch}, codebooks={NUM_CODEBOOKS}, vocab={CODEBOOK_VOCAB_SIZE}]" + ); + + let top_ids_shape = require_matching_tensor(golden, actual, AUDIO_TOP64_IDS, Dtype::I64)?; + ensure!( + top_ids_shape == [batch, NUM_CODEBOOKS, TOP_K], + "{AUDIO_TOP64_IDS} shape {top_ids_shape:?} must be [batch={batch}, codebooks={NUM_CODEBOOKS}, top_k={TOP_K}]" + ); + + let top_logprobs_shape = + require_matching_tensor(golden, actual, AUDIO_TOP64_LOGPROBS_F32, Dtype::F32)?; + ensure!( + top_logprobs_shape == [batch, NUM_CODEBOOKS, TOP_K], + "{AUDIO_TOP64_LOGPROBS_F32} shape {top_logprobs_shape:?} must be [batch={batch}, codebooks={NUM_CODEBOOKS}, top_k={TOP_K}]" + ); + + let argmax_shape = require_matching_tensor(golden, actual, AUDIO_ARGMAX_IDS, Dtype::I64)?; + ensure!( + argmax_shape == [batch, NUM_CODEBOOKS], + "{AUDIO_ARGMAX_IDS} shape {argmax_shape:?} must be [batch={batch}, codebooks={NUM_CODEBOOKS}]" + ); + + Ok(()) +} + +fn require_matching_tensor( + golden: &SafeTensors, + actual: &SafeTensors, + name: &'static str, + dtype: Dtype, +) -> Result> { + let golden = tensor(golden, name)?; + let actual = tensor(actual, name)?; + ensure!( + golden.dtype() == dtype, + "{name} golden dtype mismatch: expected {:?}, got {:?}", + dtype, + golden.dtype() + ); + ensure!( + actual.dtype() == dtype, + "{name} actual dtype mismatch: expected {:?}, got {:?}", + dtype, + actual.dtype() + ); + ensure!( + golden.shape() == actual.shape(), + "{name} shape mismatch: golden {:?} actual {:?}", + golden.shape(), + actual.shape() + ); + Ok(golden.shape().to_vec()) +} + +fn validate_prompt_surface(st: &SafeTensors, label: &str, batch: usize, seq: usize) -> Result<()> { + let lengths = i64_values(tensor(st, PROMPT_LENGTHS)?)?; + ensure!( + lengths.len() == batch, + "{label} prompt lengths element count {} must equal batch {batch}", + lengths.len() + ); + let attention = i64_values(tensor(st, PROMPT_ATTENTION_MASK)?)?; + ensure!( + attention.len() == batch * seq, + "{label} attention mask element count {} must equal batch*seq {}", + attention.len(), + batch * seq + ); + + for (row_idx, length) in lengths.iter().enumerate() { + ensure!( + *length > 0, + "{label} prompt length at row {row_idx} must be positive, got {length}" + ); + ensure!( + (*length as usize) <= seq, + "{label} prompt length at row {row_idx} exceeds seq {seq}: {length}" + ); + let row = &attention[row_idx * seq..(row_idx + 1) * seq]; + let mut mask_sum = 0i64; + for (col_idx, value) in row.iter().enumerate() { + ensure!( + *value == 0 || *value == 1, + "{label} attention mask at row {row_idx} col {col_idx} must be 0/1, got {value}" + ); + mask_sum += *value; + } + ensure!( + mask_sum == *length, + "{label} attention mask sum at row {row_idx} must equal prompt length {length}, got {mask_sum}" + ); + } + + Ok(()) +} + +fn compare_bf16( + golden: &SafeTensors, + actual: &SafeTensors, + name: &'static str, + abs_tol: f32, + mean_abs_tol: f32, +) -> Result { + let golden = tensor(golden, name)?; + let actual = tensor(actual, name)?; + ensure!(golden.dtype() == Dtype::BF16, "{name} golden must be BF16"); + ensure!(actual.dtype() == Dtype::BF16, "{name} actual must be BF16"); + compare_float_values( + name, + &bf16_values(golden)?, + &bf16_values(actual)?, + abs_tol, + mean_abs_tol, + ) +} + +fn compare_float_values( + name: &'static str, + golden: &[f32], + actual: &[f32], + abs_tol: f32, + mean_abs_tol: f32, +) -> Result { + ensure!( + golden.len() == actual.len(), + "{name} element count mismatch: golden {} actual {}", + golden.len(), + actual.len() + ); + let mut non_finite = 0usize; + let diffs: Vec = golden + .iter() + .zip(actual) + .map(|(golden, actual)| { + let diff = (*golden - *actual).abs(); + if diff.is_finite() { + diff + } else { + non_finite += 1; + f32::INFINITY + } + }) + .collect(); + let stats = stats(&diffs); + Ok(TensorComparison { + name, + elements: diffs.len(), + exact_mismatches: non_finite, + max_abs: stats.max_abs, + mean_abs: stats.mean_abs, + rmse: stats.rmse, + p99_abs: stats.p99_abs, + abs_tol, + mean_abs_tol, + passed: non_finite == 0 && stats.max_abs <= abs_tol && stats.mean_abs <= mean_abs_tol, + }) +} + +#[derive(Debug, Clone, Copy)] +struct FloatStats { + max_abs: f32, + mean_abs: f32, + rmse: f32, + p99_abs: f32, +} + +fn stats(diffs: &[f32]) -> FloatStats { + if diffs.is_empty() { + return FloatStats { + max_abs: 0.0, + mean_abs: 0.0, + rmse: 0.0, + p99_abs: 0.0, + }; + } + let mut sorted = diffs.to_vec(); + sorted.sort_by(|a, b| a.total_cmp(b)); + let sum: f64 = diffs.iter().map(|value| f64::from(*value)).sum(); + let sum_sq: f64 = diffs + .iter() + .map(|value| { + let value = f64::from(*value); + value * value + }) + .sum(); + let p99_idx = ((sorted.len() as f64 * 0.99).ceil() as usize) + .saturating_sub(1) + .min(sorted.len() - 1); + FloatStats { + max_abs: *sorted.last().expect("non-empty"), + mean_abs: (sum / diffs.len() as f64) as f32, + rmse: (sum_sq / diffs.len() as f64).sqrt() as f32, + p99_abs: sorted[p99_idx], + } +} + +fn cosine_similarity(golden: &[f32], actual: &[f32]) -> Result { + ensure!( + golden.len() == actual.len(), + "cosine element count mismatch: golden {} actual {}", + golden.len(), + actual.len() + ); + ensure!(!golden.is_empty(), "cosine requires at least one element"); + + let mut dot = 0.0f64; + let mut golden_norm_sq = 0.0f64; + let mut actual_norm_sq = 0.0f64; + for (golden, actual) in golden.iter().zip(actual) { + ensure!( + golden.is_finite() && actual.is_finite(), + "cosine inputs must be finite" + ); + let golden = f64::from(*golden); + let actual = f64::from(*actual); + dot += golden * actual; + golden_norm_sq += golden * golden; + actual_norm_sq += actual * actual; + } + ensure!( + golden_norm_sq > 0.0 && actual_norm_sq > 0.0, + "cosine inputs must have non-zero norm" + ); + Ok((dot / (golden_norm_sq.sqrt() * actual_norm_sq.sqrt())) as f32) +} + +fn max_argmax_regret(golden_logits: &[f32], actual_argmax: &[i64]) -> Result { + let rows = actual_argmax.len(); + ensure!(rows > 0, "argmax regret requires at least one row"); + ensure!( + golden_logits.len().is_multiple_of(rows), + "golden logits length {} is not divisible by argmax rows {}", + golden_logits.len(), + rows + ); + let vocab = golden_logits.len() / rows; + ensure!(vocab > 0, "argmax regret requires non-empty vocab"); + + let mut max_regret = 0.0f32; + for (row_idx, actual_id) in actual_argmax.iter().enumerate() { + ensure!( + *actual_id >= 0 && (*actual_id as usize) < vocab, + "actual argmax id {} out of range for vocab {} at row {}", + actual_id, + vocab, + row_idx + ); + let row = &golden_logits[row_idx * vocab..(row_idx + 1) * vocab]; + let golden_best = row + .iter() + .copied() + .max_by(|a, b| a.total_cmp(b)) + .context("argmax regret row is empty")?; + let actual_score = row[*actual_id as usize]; + let regret = golden_best - actual_score; + if regret > max_regret { + max_regret = regret; + } + } + Ok(max_regret) +} + +fn topk_overlap_from_logits( + golden_logits: &[f32], + actual_logits: &[f32], + rows: usize, + k: usize, +) -> Result<(usize, f32)> { + ensure!(rows > 0, "top-k overlap requires at least one row"); + ensure!(k > 0, "top-k overlap requires k > 0"); + ensure!( + golden_logits.len() == actual_logits.len(), + "top-k overlap element count mismatch: golden {} actual {}", + golden_logits.len(), + actual_logits.len() + ); + ensure!( + golden_logits.len().is_multiple_of(rows), + "top-k overlap logits length {} is not divisible by rows {}", + golden_logits.len(), + rows + ); + let vocab = golden_logits.len() / rows; + ensure!( + k <= vocab, + "top-k overlap k {} exceeds per-codebook vocab {}", + k, + vocab + ); + + let mut min_overlap = usize::MAX; + let mut overlap_sum = 0usize; + for row_idx in 0..rows { + let start = row_idx * vocab; + let end = start + vocab; + let golden_top = topk_indices(&golden_logits[start..end], k)?; + let actual_top = topk_indices(&actual_logits[start..end], k)?; + let overlap = actual_top + .iter() + .filter(|idx| golden_top.contains(idx)) + .count(); + min_overlap = min_overlap.min(overlap); + overlap_sum += overlap; + } + + Ok((min_overlap, overlap_sum as f32 / rows as f32)) +} + +fn topk_indices(values: &[f32], k: usize) -> Result> { + ensure!( + k <= values.len(), + "top-k k {} exceeds row length {}", + k, + values.len() + ); + let mut indexed: Vec<_> = values.iter().copied().enumerate().collect(); + for (idx, value) in &indexed { + ensure!( + value.is_finite(), + "top-k value at index {idx} must be finite" + ); + } + indexed.sort_by(|(left_idx, left), (right_idx, right)| { + right.total_cmp(left).then_with(|| left_idx.cmp(right_idx)) + }); + indexed.truncate(k); + Ok(indexed.into_iter().map(|(idx, _)| idx).collect()) +} + +fn tensor<'a>(st: &'a SafeTensors, name: &str) -> Result> { + st.tensor(name) + .with_context(|| format!("safetensors missing tensor {name}")) +} + +fn i64_values(tensor: TensorView<'_>) -> Result> { + bytes_to_chunks(tensor.data(), 8, "i64")?; + Ok(tensor + .data() + .chunks_exact(8) + .map(|bytes| { + i64::from_le_bytes([ + bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7], + ]) + }) + .collect()) +} + +fn f32_values(tensor: TensorView<'_>) -> Result> { + bytes_to_chunks(tensor.data(), 4, "f32")?; + Ok(tensor + .data() + .chunks_exact(4) + .map(|bytes| f32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])) + .collect()) +} + +fn bf16_values(tensor: TensorView<'_>) -> Result> { + bytes_to_chunks(tensor.data(), 2, "bf16")?; + Ok(tensor + .data() + .chunks_exact(2) + .map(|bytes| bf16::from_bits(u16::from_le_bytes([bytes[0], bytes[1]])).to_f32()) + .collect()) +} + +fn bytes_to_chunks(bytes: &[u8], chunk: usize, dtype: &str) -> Result<()> { + ensure!( + bytes.len().is_multiple_of(chunk), + "{dtype} tensor byte length {} is not divisible by {chunk}", + bytes.len() + ); + Ok(()) +} + +#[cfg(test)] +mod tests { + use std::borrow::Cow; + use std::collections::BTreeMap; + + use safetensors::tensor::View; + + use super::*; + + const GOLDEN: &str = concat!( + env!("CARGO_MANIFEST_DIR"), + "/../test_data/higgs-one-step-audio-logits.safetensors" + ); + + #[test] + fn golden_compares_equal_to_itself() { + let comparison = + compare_one_step_files(GOLDEN, GOLDEN, OneStepTolerances::default()).unwrap(); + assert!(comparison.passed()); + assert_eq!(comparison.tensors.len(), 8); + for tensor in comparison.tensors { + assert_eq!(tensor.max_abs, 0.0, "{}", tensor.name); + assert_eq!(tensor.mean_abs, 0.0, "{}", tensor.name); + } + } + + #[test] + fn semantic_comparator_accepts_golden_self_comparison() { + let comparison = + compare_one_step_semantic_files(GOLDEN, GOLDEN, OneStepSemanticTolerances::default()) + .unwrap(); + assert!(comparison.passed()); + assert!(comparison.prompt_exact); + assert!(comparison.audio_argmax_exact); + assert_eq!(comparison.hidden_cosine, 1.0); + assert_eq!(comparison.logits_cosine, 1.0); + assert_eq!(comparison.max_argmax_regret, 0.0); + assert_eq!(comparison.top64_min_overlap, 64); + } + + #[test] + fn comparison_accepts_dynamic_batch_schema() { + let bytes = dynamic_one_step_bytes(2, 3); + let st = SafeTensors::deserialize(&bytes).unwrap(); + + let strict = compare_one_step_safetensors(&st, &st, OneStepTolerances::default()).unwrap(); + assert!(strict.passed()); + + let semantic = + compare_one_step_semantic_safetensors(&st, &st, OneStepSemanticTolerances::default()) + .unwrap(); + assert!(semantic.passed()); + assert_eq!(semantic.top64_min_overlap, TOP_K); + } + + #[test] + fn comparison_rejects_prompt_length_out_of_range() { + let tmp = tempfile::NamedTempFile::new().unwrap(); + let mut bytes = std::fs::read(GOLDEN).unwrap(); + let first_length = tensor_data_offset(&bytes, PROMPT_LENGTHS); + bytes[first_length..first_length + 8].copy_from_slice(&11i64.to_le_bytes()); + std::fs::write(tmp.path(), bytes).unwrap(); + + let err = compare_one_step_files(GOLDEN, tmp.path(), OneStepTolerances::default()) + .unwrap_err() + .to_string(); + assert!(err.contains("actual prompt length at row 0 exceeds seq 10")); + } + + #[test] + fn comparison_rejects_attention_mask_length_mismatch() { + let tmp = tempfile::NamedTempFile::new().unwrap(); + let mut bytes = std::fs::read(GOLDEN).unwrap(); + let first_mask = tensor_data_offset(&bytes, PROMPT_ATTENTION_MASK); + bytes[first_mask..first_mask + 8].copy_from_slice(&0i64.to_le_bytes()); + std::fs::write(tmp.path(), bytes).unwrap(); + + let err = compare_one_step_files(GOLDEN, tmp.path(), OneStepTolerances::default()) + .unwrap_err() + .to_string(); + assert!(err.contains("actual attention mask sum at row 0")); + } + + #[test] + fn semantic_comparator_rejects_argmax_mismatch() { + let tmp = tempfile::NamedTempFile::new().unwrap(); + let mut bytes = std::fs::read(GOLDEN).unwrap(); + let first_argmax = tensor_data_offset(&bytes, AUDIO_ARGMAX_IDS); + let original = + i64::from_le_bytes(bytes[first_argmax..first_argmax + 8].try_into().unwrap()); + bytes[first_argmax..first_argmax + 8].copy_from_slice(&(original + 1).to_le_bytes()); + std::fs::write(tmp.path(), bytes).unwrap(); + + let comparison = compare_one_step_semantic_files( + GOLDEN, + tmp.path(), + OneStepSemanticTolerances::default(), + ) + .unwrap(); + assert!(!comparison.audio_argmax_exact); + assert!(!comparison.passed()); + let err = ensure_semantic_comparison_passed(&comparison) + .unwrap_err() + .to_string(); + assert!(err.contains("argmax_exact=false")); + } + + #[test] + fn semantic_comparator_rejects_prompt_mismatch() { + let tmp = tempfile::NamedTempFile::new().unwrap(); + let mut bytes = std::fs::read(GOLDEN).unwrap(); + let first_prompt_id = tensor_data_offset(&bytes, PROMPT_INPUT_IDS); + let original = i64::from_le_bytes( + bytes[first_prompt_id..first_prompt_id + 8] + .try_into() + .unwrap(), + ); + bytes[first_prompt_id..first_prompt_id + 8].copy_from_slice(&(original + 1).to_le_bytes()); + std::fs::write(tmp.path(), bytes).unwrap(); + + let comparison = compare_one_step_semantic_files( + GOLDEN, + tmp.path(), + OneStepSemanticTolerances::default(), + ) + .unwrap(); + assert!(!comparison.prompt_exact); + assert!(comparison.audio_argmax_exact); + assert!(!comparison.passed()); + let err = ensure_semantic_comparison_passed(&comparison) + .unwrap_err() + .to_string(); + assert!(err.contains("prompt_exact=false")); + assert!(err.contains("top64_mean_overlap=")); + } + + #[test] + fn topk_overlap_reports_min_and_mean_by_codebook() { + let mut golden = Vec::new(); + let mut actual = Vec::new(); + for row in 0..8 { + for col in 0..8 { + golden.push((8 - col) as f32); + actual.push(if row == 0 { + col as f32 + } else { + (8 - col) as f32 + }); + } + } + let (min_overlap, mean_overlap) = topk_overlap_from_logits(&golden, &actual, 8, 4).unwrap(); + assert_eq!(min_overlap, 0); + assert_eq!(mean_overlap, 3.5); + } + + #[test] + fn topk_overlap_uses_dynamic_row_count() { + let mut golden = Vec::new(); + let mut actual = Vec::new(); + for row in 0..16 { + for col in 0..8 { + golden.push((8 - col) as f32); + actual.push(if row < 2 { + col as f32 + } else { + (8 - col) as f32 + }); + } + } + let (min_overlap, mean_overlap) = + topk_overlap_from_logits(&golden, &actual, 16, 4).unwrap(); + assert_eq!(min_overlap, 0); + assert_eq!(mean_overlap, 3.5); + } + + #[test] + fn comparator_rejects_logit_drift_beyond_tolerance() { + let tmp = tempfile::NamedTempFile::new().unwrap(); + let mut bytes = std::fs::read(GOLDEN).unwrap(); + let first_logit = tensor_data_offset(&bytes, AUDIO_LOGITS_F32); + bytes[first_logit + 3] ^= 0x40; + std::fs::write(tmp.path(), bytes).unwrap(); + + let comparison = + compare_one_step_files(GOLDEN, tmp.path(), OneStepTolerances::default()).unwrap(); + let logits = comparison + .tensors + .iter() + .find(|tensor| tensor.name == AUDIO_LOGITS_F32) + .unwrap(); + assert!(!logits.passed); + assert!(!comparison.passed()); + let err = ensure_comparison_passed(&comparison) + .unwrap_err() + .to_string(); + assert!(err.contains(AUDIO_LOGITS_F32)); + } + + fn tensor_data_offset(bytes: &[u8], name: &str) -> usize { + let header_len = u64::from_le_bytes(bytes[..8].try_into().unwrap()) as usize; + let header: serde_json::Value = serde_json::from_slice(&bytes[8..8 + header_len]).unwrap(); + let start = header[name]["data_offsets"][0].as_u64().unwrap() as usize; + 8 + header_len + start + } + + fn dynamic_one_step_bytes(batch: usize, seq: usize) -> Vec { + let rows = batch * NUM_CODEBOOKS; + let mut prompt = Vec::with_capacity(batch * seq); + let mut mask = Vec::with_capacity(batch * seq); + for row in 0..batch { + for col in 0..seq { + prompt.push((row * seq + col + 1) as i64); + mask.push(1); + } + } + let lengths = vec![seq as i64; batch]; + + let hidden: Vec<_> = (0..batch * HIDDEN_SIZE) + .map(|idx| bf16::from_f32((idx % 17 + 1) as f32 / 17.0)) + .collect(); + let mut logits = vec![0.0f32; rows * CODEBOOK_VOCAB_SIZE]; + let mut argmax = Vec::with_capacity(rows); + let mut top_ids = Vec::with_capacity(rows * TOP_K); + let mut top_logprobs = Vec::with_capacity(rows * TOP_K); + for row in 0..rows { + let best = row % CODEBOOK_VOCAB_SIZE; + logits[row * CODEBOOK_VOCAB_SIZE + best] = 1.0; + argmax.push(best as i64); + for offset in 0..TOP_K { + top_ids.push(((best + offset) % CODEBOOK_VOCAB_SIZE) as i64); + top_logprobs.push(-(offset as f32)); + } + } + + let tensors = BTreeMap::from([ + ( + PROMPT_INPUT_IDS.to_string(), + test_i64(&[batch, seq], &prompt), + ), + ( + PROMPT_ATTENTION_MASK.to_string(), + test_i64(&[batch, seq], &mask), + ), + (PROMPT_LENGTHS.to_string(), test_i64(&[batch], &lengths)), + ( + FINAL_HIDDEN_BF16.to_string(), + test_bf16(&[batch, HIDDEN_SIZE], &hidden), + ), + ( + AUDIO_LOGITS_F32.to_string(), + test_f32(&[batch, NUM_CODEBOOKS, CODEBOOK_VOCAB_SIZE], &logits), + ), + ( + AUDIO_TOP64_IDS.to_string(), + test_i64(&[batch, NUM_CODEBOOKS, TOP_K], &top_ids), + ), + ( + AUDIO_TOP64_LOGPROBS_F32.to_string(), + test_f32(&[batch, NUM_CODEBOOKS, TOP_K], &top_logprobs), + ), + ( + AUDIO_ARGMAX_IDS.to_string(), + test_i64(&[batch, NUM_CODEBOOKS], &argmax), + ), + ]); + safetensors::serialize(tensors, None).unwrap() + } + + #[derive(Clone)] + struct TestTensor { + dtype: Dtype, + shape: Vec, + data: Vec, + } + + impl View for TestTensor { + fn dtype(&self) -> Dtype { + self.dtype + } + + fn shape(&self) -> &[usize] { + &self.shape + } + + fn data(&self) -> Cow<'_, [u8]> { + Cow::Borrowed(&self.data) + } + + fn data_len(&self) -> usize { + self.data.len() + } + } + + fn test_i64(shape: &[usize], values: &[i64]) -> TestTensor { + TestTensor { + dtype: Dtype::I64, + shape: shape.to_vec(), + data: values + .iter() + .flat_map(|value| value.to_le_bytes()) + .collect(), + } + } + + fn test_f32(shape: &[usize], values: &[f32]) -> TestTensor { + TestTensor { + dtype: Dtype::F32, + shape: shape.to_vec(), + data: values + .iter() + .flat_map(|value| value.to_le_bytes()) + .collect(), + } + } + + fn test_bf16(shape: &[usize], values: &[bf16]) -> TestTensor { + TestTensor { + dtype: Dtype::BF16, + shape: shape.to_vec(), + data: values + .iter() + .flat_map(|value| value.to_bits().to_le_bytes()) + .collect(), + } + } +} diff --git a/pegainfer-higgs-audio/src/config.rs b/pegainfer-higgs-audio/src/config.rs new file mode 100644 index 000000000..760333e51 --- /dev/null +++ b/pegainfer-higgs-audio/src/config.rs @@ -0,0 +1,269 @@ +use std::path::Path; + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use serde_json::Value; + +use crate::one_step_golden::CODEBOOK_VOCAB_SIZE; +use crate::one_step_golden::HIDDEN_SIZE; +use crate::one_step_golden::NUM_CODEBOOKS; + +pub const EXPECTED_MODEL_TYPE: &str = "higgs_multimodal_qwen3"; +pub const EXPECTED_ARCHITECTURE: &str = "HiggsMultimodalQwen3ForConditionalGeneration"; +pub const EXPECTED_AUDIO_ENCODER_TYPE: &str = "discrete"; +pub const EXPECTED_NUM_LAYERS: usize = 36; +pub const EXPECTED_NUM_ATTENTION_HEADS: usize = 32; +pub const EXPECTED_NUM_KV_HEADS: usize = 8; +pub const EXPECTED_HEAD_DIM: usize = 128; +pub const EXPECTED_INTERMEDIATE_SIZE: usize = 9728; +pub const EXPECTED_TEXT_VOCAB_SIZE: usize = 151_936; +pub const EXPECTED_ROPE_THETA: u64 = 1_000_000; +pub const EXPECTED_MODEL_CARD_CONTEXT: usize = 8192; + +#[derive(Debug, Clone, PartialEq)] +pub struct HiggsConfig { + pub model_type: String, + pub architecture: String, + pub audio_token_id: i64, + pub text: TextConfig, + pub audio: AudioEncoderConfig, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct TextConfig { + pub hidden_size: usize, + pub intermediate_size: usize, + pub num_hidden_layers: usize, + pub num_attention_heads: usize, + pub num_key_value_heads: usize, + pub head_dim: usize, + pub vocab_size: usize, + pub rms_norm_eps: f32, + pub max_position_embeddings: usize, + pub eos_token_id: u32, + pub tie_word_embeddings: bool, + pub rope_theta: u64, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct AudioEncoderConfig { + pub encoder_type: String, + pub num_codebooks: usize, + pub vocab_size: usize, + pub out_dim: usize, + pub tie_word_embeddings: bool, + pub use_delay_pattern: bool, +} + +impl HiggsConfig { + pub fn from_model_dir(model_dir: impl AsRef) -> Result { + let path = model_dir.as_ref().join("config.json"); + let value: Value = serde_json::from_slice( + &std::fs::read(&path).with_context(|| format!("read {}", path.display()))?, + ) + .with_context(|| format!("parse {}", path.display()))?; + Self::from_json(&value) + } + + pub fn from_json(value: &Value) -> Result { + let text = value + .get("text_config") + .context("Higgs config missing text_config")?; + let audio = value + .get("audio_encoder_config") + .context("Higgs config missing audio_encoder_config")?; + let config = Self { + model_type: str_field(value, "model_type")?.to_string(), + architecture: first_architecture(value)?, + audio_token_id: i64_field(value, "audio_token_id")?, + text: TextConfig { + hidden_size: usize_field(text, "hidden_size")?, + intermediate_size: usize_field(text, "intermediate_size")?, + num_hidden_layers: usize_field(text, "num_hidden_layers")?, + num_attention_heads: usize_field(text, "num_attention_heads")?, + num_key_value_heads: usize_field(text, "num_key_value_heads")?, + head_dim: usize_field(text, "head_dim")?, + vocab_size: usize_field(text, "vocab_size")?, + rms_norm_eps: f32_field(text, "rms_norm_eps")?, + max_position_embeddings: usize_field(text, "max_position_embeddings")?, + eos_token_id: u32_field(text, "eos_token_id")?, + tie_word_embeddings: bool_field(text, "tie_word_embeddings")?, + rope_theta: text + .get("rope_parameters") + .and_then(|rope| rope.get("rope_theta")) + .and_then(Value::as_u64) + .context("text_config.rope_parameters.rope_theta missing or not u64")?, + }, + audio: AudioEncoderConfig { + encoder_type: str_field(audio, "encoder_type")?.to_string(), + num_codebooks: usize_field(audio, "num_codebooks")?, + vocab_size: usize_field(audio, "vocab_size")?, + out_dim: usize_field(audio, "out_dim")?, + tie_word_embeddings: bool_field(audio, "tie_word_embeddings")?, + use_delay_pattern: bool_field(audio, "use_delay_pattern")?, + }, + }; + config.validate_current_contract()?; + Ok(config) + } + + pub fn validate_current_contract(&self) -> Result<()> { + if self.model_type != EXPECTED_MODEL_TYPE { + bail!("unexpected Higgs model_type {}", self.model_type); + } + if self.architecture != EXPECTED_ARCHITECTURE { + bail!("unexpected Higgs architecture {}", self.architecture); + } + if self.audio_token_id != -100 { + bail!( + "Higgs audio_token_id must be -100, got {}", + self.audio_token_id + ); + } + if self.text.hidden_size != HIDDEN_SIZE + || self.text.intermediate_size != EXPECTED_INTERMEDIATE_SIZE + || self.text.num_hidden_layers != EXPECTED_NUM_LAYERS + || self.text.num_attention_heads != EXPECTED_NUM_ATTENTION_HEADS + || self.text.num_key_value_heads != EXPECTED_NUM_KV_HEADS + || self.text.head_dim != EXPECTED_HEAD_DIM + || self.text.vocab_size != EXPECTED_TEXT_VOCAB_SIZE + || self.text.rope_theta != EXPECTED_ROPE_THETA + || !self.text.tie_word_embeddings + { + bail!("Higgs text_config does not match the pinned one-step contract: {self:?}"); + } + if self.audio.encoder_type != EXPECTED_AUDIO_ENCODER_TYPE + || self.audio.num_codebooks != NUM_CODEBOOKS + || self.audio.vocab_size != CODEBOOK_VOCAB_SIZE + || self.audio.out_dim != HIDDEN_SIZE + || !self.audio.tie_word_embeddings + || !self.audio.use_delay_pattern + { + bail!( + "Higgs audio_encoder_config does not match the pinned one-step contract: {self:?}" + ); + } + Ok(()) + } + + pub fn kv_bytes_per_position_bf16(&self) -> usize { + self.text.num_hidden_layers + * 2 + * self.text.num_key_value_heads + * self.text.head_dim + * std::mem::size_of::() + } + + pub fn kv_bytes_for_positions_bf16(&self, positions: usize) -> usize { + self.kv_bytes_per_position_bf16() * positions + } +} + +fn first_architecture(value: &Value) -> Result { + value + .get("architectures") + .and_then(Value::as_array) + .and_then(|items| items.first()) + .and_then(Value::as_str) + .map(str::to_string) + .context("Higgs config missing architectures[0]") +} + +fn str_field<'a>(value: &'a Value, key: &str) -> Result<&'a str> { + value + .get(key) + .and_then(Value::as_str) + .with_context(|| format!("missing string field {key}")) +} + +fn usize_field(value: &Value, key: &str) -> Result { + let raw = value + .get(key) + .and_then(Value::as_u64) + .with_context(|| format!("missing usize field {key}"))?; + usize::try_from(raw).with_context(|| format!("{key} does not fit usize")) +} + +fn u32_field(value: &Value, key: &str) -> Result { + let raw = value + .get(key) + .and_then(Value::as_u64) + .with_context(|| format!("missing u32 field {key}"))?; + u32::try_from(raw).with_context(|| format!("{key} does not fit u32")) +} + +fn i64_field(value: &Value, key: &str) -> Result { + value + .get(key) + .and_then(Value::as_i64) + .with_context(|| format!("missing i64 field {key}")) +} + +fn f32_field(value: &Value, key: &str) -> Result { + value + .get(key) + .and_then(Value::as_f64) + .map(|value| value as f32) + .with_context(|| format!("missing f32 field {key}")) +} + +fn bool_field(value: &Value, key: &str) -> Result { + value + .get(key) + .and_then(Value::as_bool) + .with_context(|| format!("missing bool field {key}")) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn minimal_config() -> serde_json::Value { + serde_json::json!({ + "architectures": [EXPECTED_ARCHITECTURE], + "audio_token_id": -100, + "model_type": EXPECTED_MODEL_TYPE, + "text_config": { + "hidden_size": HIDDEN_SIZE, + "intermediate_size": EXPECTED_INTERMEDIATE_SIZE, + "num_hidden_layers": EXPECTED_NUM_LAYERS, + "num_attention_heads": EXPECTED_NUM_ATTENTION_HEADS, + "num_key_value_heads": EXPECTED_NUM_KV_HEADS, + "head_dim": EXPECTED_HEAD_DIM, + "vocab_size": EXPECTED_TEXT_VOCAB_SIZE, + "rms_norm_eps": 1e-6, + "max_position_embeddings": 32768, + "eos_token_id": 151643, + "tie_word_embeddings": true, + "rope_parameters": {"rope_theta": EXPECTED_ROPE_THETA} + }, + "audio_encoder_config": { + "encoder_type": EXPECTED_AUDIO_ENCODER_TYPE, + "num_codebooks": NUM_CODEBOOKS, + "vocab_size": CODEBOOK_VOCAB_SIZE, + "out_dim": HIDDEN_SIZE, + "tie_word_embeddings": true, + "use_delay_pattern": true + } + }) + } + + #[test] + fn parses_pinned_higgs_config_shape_contract() { + let config = HiggsConfig::from_json(&minimal_config()).unwrap(); + assert_eq!(config.kv_bytes_per_position_bf16(), 144 * 1024); + assert_eq!( + config.kv_bytes_for_positions_bf16(EXPECTED_MODEL_CARD_CONTEXT), + 1152 * 1024 * 1024 + ); + } + + #[test] + fn rejects_non_discrete_audio_encoder() { + let mut value = minimal_config(); + value["audio_encoder_config"]["encoder_type"] = serde_json::json!("whisper"); + let err = HiggsConfig::from_json(&value).unwrap_err().to_string(); + assert!(err.contains("audio_encoder_config")); + } +} diff --git a/pegainfer-higgs-audio/src/kernel_plan.rs b/pegainfer-higgs-audio/src/kernel_plan.rs new file mode 100644 index 000000000..f17e915dd --- /dev/null +++ b/pegainfer-higgs-audio/src/kernel_plan.rs @@ -0,0 +1,141 @@ +pub struct KernelPlan { + pub model: &'static str, + pub phases: &'static [KernelPhase], +} + +pub struct KernelPhase { + pub name: &'static str, + pub ops: &'static [KernelOp], +} + +pub struct KernelOp { + pub id: &'static str, + pub rust: &'static str, + pub backend: &'static str, + pub notes: &'static str, +} + +pub static KERNEL_PLAN: KernelPlan = KernelPlan { + model: "higgs-audio", + phases: &[ + KernelPhase { + name: "artifact", + ops: &[ + KernelOp { + id: "checkpoint_header_gate", + rust: "weights::HiggsWeightManifest::from_model_dir", + backend: "safetensors header", + notes: "validates Higgs checkpoint tensor names, dtypes, and shapes without reading payloads", + }, + KernelOp { + id: "qwen3_alias_plan", + rust: "load_plan::HiggsRuntimeLoadPlan::qwen3_tensor_aliases", + backend: "metadata", + notes: "maps Higgs body.* tensors onto Qwen3 requested tensor names without a 7.5 GiB payload copy", + }, + ], + }, + KernelPhase { + name: "prefill", + ops: &[ + KernelOp { + id: "qwen3_body_prefill", + rust: "runtime_bridge::HiggsAudioRuntime::prefill_audio_from_prompt_ids -> Qwen3Executor::prefill_last_hidden_bf16", + backend: "Qwen3 runtime: CUDA + cuBLAS + FlashInfer", + notes: "runs the Higgs text/body checkpoint through the existing Qwen3 prefill path via tensor-name aliases", + }, + KernelOp { + id: "qwen3_prompt_session_prefill", + rust: "runtime_bridge::HiggsAudioRuntime::prefill_prompt_session -> Qwen3Executor::prefill_last_hidden_bf16_retained_prompt", + backend: "Qwen3 runtime: CUDA + cuBLAS + FlashInfer + paged KV", + notes: "retains prompt KV under a Higgs-owned session handle without registering a generated text token", + }, + KernelOp { + id: "fused_audio_head", + rust: "one_step_actual::compute_one_step_audio_prediction_gpu_bf16 -> ops::linear", + backend: "CUDA bf16 linear", + notes: "projects the final hidden state with tied.embedding.modality_embeddings.0.embedding.weight into 8x1026 audio logits", + }, + KernelOp { + id: "audio_topk_argmax", + rust: "one_step_actual::audio_topk_and_argmax", + backend: "CPU", + notes: "diagnostic one-step gate extracts top-64 and argmax ids from the fused audio logits", + }, + ], + }, + KernelPhase { + name: "golden", + ops: &[ + KernelOp { + id: "strict_comparison", + rust: "compare::compare_one_step_files", + backend: "CPU", + notes: "exact prompt/argmax checks plus absolute drift diagnostics for hidden, logits, and top-64 logprobs", + }, + KernelOp { + id: "semantic_comparison", + rust: "compare::compare_one_step_semantic_files", + backend: "CPU", + notes: "runtime bring-up gate using prompt exactness, argmax exactness, cosine, regret, and top-64 overlap", + }, + ], + }, + ], +}; + +pub fn kernel_plan() -> &'static KernelPlan { + &KERNEL_PLAN +} + +#[cfg(test)] +mod tests { + use super::kernel_plan; + + #[test] + fn higgs_kernel_plan_names_current_phases() { + let phase_names: Vec<_> = kernel_plan() + .phases + .iter() + .map(|phase| phase.name) + .collect(); + assert_eq!(phase_names, ["artifact", "prefill", "golden"]); + } + + #[test] + fn higgs_kernel_plan_records_runtime_backends() { + let ops: Vec<_> = kernel_plan() + .phases + .iter() + .flat_map(|phase| phase.ops.iter()) + .collect(); + + assert!(ops.iter().any(|op| { + op.id == "qwen3_body_prefill" + && op.backend == "Qwen3 runtime: CUDA + cuBLAS + FlashInfer" + })); + assert!(ops.iter().any(|op| { + op.id == "qwen3_prompt_session_prefill" && op.notes.contains("retains prompt KV") + })); + assert!( + ops.iter() + .any(|op| op.id == "fused_audio_head" && op.backend == "CUDA bf16 linear") + ); + assert!( + ops.iter() + .any(|op| op.id == "semantic_comparison" && op.backend == "CPU") + ); + } + + #[test] + fn higgs_kernel_plan_keeps_alias_copy_boundary_visible() { + let alias_op = kernel_plan() + .phases + .iter() + .flat_map(|phase| phase.ops.iter()) + .find(|op| op.id == "qwen3_alias_plan") + .expect("qwen3 alias op should be in the plan"); + + assert!(alias_op.notes.contains("without a 7.5 GiB payload copy")); + } +} diff --git a/pegainfer-higgs-audio/src/layer_dump.rs b/pegainfer-higgs-audio/src/layer_dump.rs new file mode 100644 index 000000000..9dd332fdb --- /dev/null +++ b/pegainfer-higgs-audio/src/layer_dump.rs @@ -0,0 +1,166 @@ +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::path::Path; +use std::path::PathBuf; + +use anyhow::Context; +use anyhow::Result; +use anyhow::ensure; +use half::bf16; + +use crate::compare::FINAL_HIDDEN_BF16; +use crate::compare::PROMPT_ATTENTION_MASK; +use crate::compare::PROMPT_INPUT_IDS; +use crate::compare::PROMPT_LENGTHS; +use crate::one_step_actual::PromptTensors; +use crate::one_step_actual::owned_bf16; +use crate::one_step_actual::owned_i64; +use crate::one_step_golden::HIDDEN_SIZE; + +pub const NUM_LAYERS: usize = 36; +pub const EMBEDDING_HIDDEN_BF16: &str = "embedding.last_hidden.bf16"; + +#[derive(Debug, Clone, PartialEq)] +pub struct LayerHiddenDumpSummary { + pub output_path: PathBuf, + pub prompt_tokens: usize, + pub layers: usize, + pub hidden_values_per_layer: usize, +} + +pub fn layer_hidden_tensor_name(layer_idx: usize) -> String { + format!("layer.{layer_idx:02}.last_hidden.bf16") +} + +pub fn write_layer_hidden_dump( + output_path: impl AsRef, + prompt: &PromptTensors, + embedding_hidden: &[bf16], + layer_hidden: &[Vec], + final_normed: &[bf16], +) -> Result { + ensure!( + embedding_hidden.len() == HIDDEN_SIZE, + "embedding hidden len mismatch: expected {HIDDEN_SIZE}, got {}", + embedding_hidden.len() + ); + ensure!( + layer_hidden.len() == NUM_LAYERS, + "expected {NUM_LAYERS} layer snapshots, got {}", + layer_hidden.len() + ); + ensure!( + final_normed.len() == HIDDEN_SIZE, + "final normed hidden len mismatch: expected {HIDDEN_SIZE}, got {}", + final_normed.len() + ); + + let mut tensors = BTreeMap::from([ + ( + PROMPT_INPUT_IDS.to_string(), + owned_i64( + &[1, prompt.input_ids_padded.len()], + &prompt.input_ids_padded, + ), + ), + ( + PROMPT_ATTENTION_MASK.to_string(), + owned_i64(&[1, prompt.attention_mask.len()], &prompt.attention_mask), + ), + ( + PROMPT_LENGTHS.to_string(), + owned_i64(&[prompt.lengths.len()], &prompt.lengths), + ), + ( + EMBEDDING_HIDDEN_BF16.to_string(), + owned_bf16(&[1, HIDDEN_SIZE], embedding_hidden), + ), + ( + FINAL_HIDDEN_BF16.to_string(), + owned_bf16(&[1, HIDDEN_SIZE], final_normed), + ), + ]); + for (layer_idx, hidden) in layer_hidden.iter().enumerate() { + ensure!( + hidden.len() == HIDDEN_SIZE, + "layer {layer_idx} hidden len mismatch: expected {HIDDEN_SIZE}, got {}", + hidden.len() + ); + tensors.insert( + layer_hidden_tensor_name(layer_idx), + owned_bf16(&[1, HIDDEN_SIZE], hidden), + ); + } + + let output_path = output_path.as_ref(); + let metadata = HashMap::from([( + "fixture_kind".to_string(), + "higgs-prefill-layer-hidden-actual".to_string(), + )]); + safetensors::serialize_to_file(tensors, Some(metadata), output_path) + .with_context(|| format!("write {}", output_path.display()))?; + + Ok(LayerHiddenDumpSummary { + output_path: output_path.to_path_buf(), + prompt_tokens: prompt.prompt_ids()?.len(), + layers: layer_hidden.len(), + hidden_values_per_layer: HIDDEN_SIZE, + }) +} + +#[derive(Debug, Clone, PartialEq)] +pub struct StageDumpSummary { + pub output_path: PathBuf, + pub prompt_tokens: usize, + pub stages: usize, + pub values: usize, +} + +pub fn write_stage_dump( + output_path: impl AsRef, + prompt: &PromptTensors, + stages: &[(String, Vec)], +) -> Result { + ensure!(!stages.is_empty(), "stage dump requires at least one stage"); + let mut tensors = BTreeMap::from([ + ( + PROMPT_INPUT_IDS.to_string(), + owned_i64( + &[1, prompt.input_ids_padded.len()], + &prompt.input_ids_padded, + ), + ), + ( + PROMPT_ATTENTION_MASK.to_string(), + owned_i64(&[1, prompt.attention_mask.len()], &prompt.attention_mask), + ), + ( + PROMPT_LENGTHS.to_string(), + owned_i64(&[prompt.lengths.len()], &prompt.lengths), + ), + ]); + let mut values = 0usize; + for (name, stage_values) in stages { + ensure!(!stage_values.is_empty(), "stage {name} must not be empty"); + values += stage_values.len(); + tensors.insert( + name.clone(), + owned_bf16(&[1, stage_values.len()], stage_values), + ); + } + + let output_path = output_path.as_ref(); + let metadata = HashMap::from([( + "fixture_kind".to_string(), + "higgs-layer0-stage-actual".to_string(), + )]); + safetensors::serialize_to_file(tensors, Some(metadata), output_path) + .with_context(|| format!("write {}", output_path.display()))?; + + Ok(StageDumpSummary { + output_path: output_path.to_path_buf(), + prompt_tokens: prompt.prompt_ids()?.len(), + stages: stages.len(), + values, + }) +} diff --git a/pegainfer-higgs-audio/src/lib.rs b/pegainfer-higgs-audio/src/lib.rs new file mode 100644 index 000000000..272125c51 --- /dev/null +++ b/pegainfer-higgs-audio/src/lib.rs @@ -0,0 +1,21 @@ +//! Higgs Audio model-line scaffolding. +//! +//! This crate starts with the artifact and golden-contract boundary for the +//! zero-shot, one-step `[8, 1026]` audio-logits gate. Runtime execution is added +//! behind this boundary in later slices so stale fixture assumptions cannot leak +//! into the model implementation. + +pub mod compare; +pub mod config; +pub mod kernel_plan; +pub mod layer_dump; +pub mod load_plan; +pub mod materialize_qwen3; +pub mod one_step_actual; +pub mod one_step_golden; +#[cfg(feature = "runtime-qwen3")] +pub mod runtime_bridge; +pub mod runtime_source; +pub mod weights; + +pub use kernel_plan::kernel_plan; diff --git a/pegainfer-higgs-audio/src/load_plan.rs b/pegainfer-higgs-audio/src/load_plan.rs new file mode 100644 index 000000000..0e7308e7b --- /dev/null +++ b/pegainfer-higgs-audio/src/load_plan.rs @@ -0,0 +1,362 @@ +use std::collections::BTreeMap; +use std::collections::BTreeSet; + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; + +use crate::config::HiggsConfig; +use crate::weights::BODY_NORM; +use crate::weights::FUSED_MODALITY_EMBEDDING; +use crate::weights::HiggsWeightManifest; +use crate::weights::TEXT_EMBEDDING; +use crate::weights::TensorHeaderSpec; +use crate::weights::expected_checkpoint_tensors; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum TensorRole { + TextEmbedding, + FusedAudioHead, + BodyNorm, + LayerInputLayernorm, + LayerPostAttentionLayernorm, + LayerQProj, + LayerKProj, + LayerVProj, + LayerOProj, + LayerQNorm, + LayerKNorm, + LayerGateProj, + LayerUpProj, + LayerDownProj, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PlannedTensor { + pub checkpoint_name: String, + pub shard_file: String, + pub role: TensorRole, + pub loader_slot: String, + pub dtype: &'static str, + pub shape: Vec, + pub elements: usize, + pub bytes: usize, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct LoadPlanSummary { + pub tensors: usize, + pub shard_files: usize, + pub bf16_bytes: usize, + pub body_tensors: usize, + pub qwen3_backbone_tensors: usize, + pub higgs_head_tensors: usize, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct HiggsRuntimeLoadPlan { + pub tensors: Vec, +} + +impl HiggsRuntimeLoadPlan { + pub fn from_manifest(config: &HiggsConfig, manifest: &HiggsWeightManifest) -> Result { + manifest.validate_for_config(config)?; + let tensors = expected_checkpoint_tensors(config) + .into_iter() + .map(|spec| planned_tensor(spec, manifest)) + .collect::>>()?; + Ok(Self { tensors }) + } + + pub fn summary(&self) -> LoadPlanSummary { + let shard_files: BTreeSet<_> = self + .tensors + .iter() + .map(|tensor| tensor.shard_file.as_str()) + .collect(); + LoadPlanSummary { + tensors: self.tensors.len(), + shard_files: shard_files.len(), + bf16_bytes: self + .tensors + .iter() + .filter(|tensor| tensor.dtype == "BF16") + .map(|tensor| tensor.bytes) + .sum(), + body_tensors: self + .tensors + .iter() + .filter(|tensor| tensor.checkpoint_name.starts_with("body.")) + .count(), + qwen3_backbone_tensors: self + .tensors + .iter() + .filter(|tensor| tensor.loader_slot.starts_with("qwen3.")) + .count(), + higgs_head_tensors: self + .tensors + .iter() + .filter(|tensor| tensor.loader_slot.starts_with("higgs.")) + .count(), + } + } + + pub fn tensor(&self, checkpoint_name: &str) -> Option<&PlannedTensor> { + self.tensors + .iter() + .find(|tensor| tensor.checkpoint_name == checkpoint_name) + } + + pub fn qwen3_tensor_aliases(&self) -> Result> { + self.tensors + .iter() + .filter(|tensor| tensor.loader_slot.starts_with("qwen3.")) + .map(|tensor| { + Ok(( + qwen3_tensor_name(&tensor.loader_slot)?, + tensor.checkpoint_name.clone(), + )) + }) + .collect() + } +} + +fn planned_tensor(spec: TensorHeaderSpec, manifest: &HiggsWeightManifest) -> Result { + let role = tensor_role(&spec.name)?; + let loader_slot = loader_slot(&spec.name, role)?; + let shard_file = manifest + .weight_map + .get(&spec.name) + .with_context(|| format!("manifest missing tensor {}", spec.name))? + .clone(); + let elements = spec.shape.iter().product::(); + let bytes_per_element = dtype_bytes(spec.dtype)?; + Ok(PlannedTensor { + checkpoint_name: spec.name, + shard_file, + role, + loader_slot, + dtype: spec.dtype, + shape: spec.shape, + elements, + bytes: elements * bytes_per_element, + }) +} + +fn tensor_role(name: &str) -> Result { + if name == TEXT_EMBEDDING { + return Ok(TensorRole::TextEmbedding); + } + if name == FUSED_MODALITY_EMBEDDING { + return Ok(TensorRole::FusedAudioHead); + } + if name == BODY_NORM { + return Ok(TensorRole::BodyNorm); + } + let suffix = layer_suffix(name)?; + match suffix { + "input_layernorm.weight" => Ok(TensorRole::LayerInputLayernorm), + "post_attention_layernorm.weight" => Ok(TensorRole::LayerPostAttentionLayernorm), + "self_attn.q_proj.weight" => Ok(TensorRole::LayerQProj), + "self_attn.k_proj.weight" => Ok(TensorRole::LayerKProj), + "self_attn.v_proj.weight" => Ok(TensorRole::LayerVProj), + "self_attn.o_proj.weight" => Ok(TensorRole::LayerOProj), + "self_attn.q_norm.weight" => Ok(TensorRole::LayerQNorm), + "self_attn.k_norm.weight" => Ok(TensorRole::LayerKNorm), + "mlp.gate_proj.weight" => Ok(TensorRole::LayerGateProj), + "mlp.up_proj.weight" => Ok(TensorRole::LayerUpProj), + "mlp.down_proj.weight" => Ok(TensorRole::LayerDownProj), + _ => bail!("unsupported Higgs layer tensor suffix {suffix} in {name}"), + } +} + +fn loader_slot(name: &str, role: TensorRole) -> Result { + match role { + TensorRole::TextEmbedding => Ok("qwen3.embed_tokens".to_string()), + TensorRole::FusedAudioHead => Ok("higgs.fused_audio_head".to_string()), + TensorRole::BodyNorm => Ok("qwen3.norm".to_string()), + _ => { + let layer = layer_index(name)?; + Ok(format!("qwen3.layers.{layer}.{}", layer_suffix(name)?)) + } + } +} + +pub fn qwen3_tensor_name(loader_slot: &str) -> Result { + match loader_slot { + "qwen3.embed_tokens" => Ok("model.embed_tokens.weight".to_string()), + "qwen3.norm" => Ok("model.norm.weight".to_string()), + slot if slot.starts_with("qwen3.layers.") => { + Ok(format!("model.layers.{}", &slot["qwen3.layers.".len()..])) + } + _ => bail!("loader slot {loader_slot} is not part of the Qwen3 body view"), + } +} + +fn layer_index(name: &str) -> Result { + let rest = name + .strip_prefix("body.layers.") + .with_context(|| format!("{name} is not a body layer tensor"))?; + let Some((layer, _suffix)) = rest.split_once('.') else { + bail!("{name} is missing layer suffix"); + }; + layer + .parse::() + .with_context(|| format!("{name} has invalid layer index")) +} + +fn layer_suffix(name: &str) -> Result<&str> { + let rest = name + .strip_prefix("body.layers.") + .with_context(|| format!("{name} is not a body layer tensor"))?; + let Some((_layer, suffix)) = rest.split_once('.') else { + bail!("{name} is missing layer suffix"); + }; + Ok(suffix) +} + +fn dtype_bytes(dtype: &str) -> Result { + match dtype { + "BF16" => Ok(2), + "F32" | "I32" => Ok(4), + "I64" => Ok(8), + _ => bail!("unsupported dtype in Higgs runtime load plan: {dtype}"), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::EXPECTED_ARCHITECTURE; + use crate::config::EXPECTED_AUDIO_ENCODER_TYPE; + use crate::config::EXPECTED_HEAD_DIM; + use crate::config::EXPECTED_INTERMEDIATE_SIZE; + use crate::config::EXPECTED_MODEL_TYPE; + use crate::config::EXPECTED_NUM_ATTENTION_HEADS; + use crate::config::EXPECTED_NUM_KV_HEADS; + use crate::config::EXPECTED_NUM_LAYERS; + use crate::config::EXPECTED_ROPE_THETA; + use crate::config::EXPECTED_TEXT_VOCAB_SIZE; + use crate::one_step_golden::CODEBOOK_VOCAB_SIZE; + use crate::one_step_golden::HIDDEN_SIZE; + use crate::one_step_golden::NUM_CODEBOOKS; + use crate::weights::fused_modality_shape; + use crate::weights::required_body_tensors; + + fn config() -> HiggsConfig { + HiggsConfig::from_json(&serde_json::json!({ + "architectures": [EXPECTED_ARCHITECTURE], + "audio_token_id": -100, + "model_type": EXPECTED_MODEL_TYPE, + "text_config": { + "hidden_size": HIDDEN_SIZE, + "intermediate_size": EXPECTED_INTERMEDIATE_SIZE, + "num_hidden_layers": EXPECTED_NUM_LAYERS, + "num_attention_heads": EXPECTED_NUM_ATTENTION_HEADS, + "num_key_value_heads": EXPECTED_NUM_KV_HEADS, + "head_dim": EXPECTED_HEAD_DIM, + "vocab_size": EXPECTED_TEXT_VOCAB_SIZE, + "rms_norm_eps": 1e-6, + "max_position_embeddings": 32768, + "eos_token_id": 151643, + "tie_word_embeddings": true, + "rope_parameters": {"rope_theta": EXPECTED_ROPE_THETA} + }, + "audio_encoder_config": { + "encoder_type": EXPECTED_AUDIO_ENCODER_TYPE, + "num_codebooks": NUM_CODEBOOKS, + "vocab_size": CODEBOOK_VOCAB_SIZE, + "out_dim": HIDDEN_SIZE, + "tie_word_embeddings": true, + "use_delay_pattern": true + } + })) + .unwrap() + } + + fn manifest() -> HiggsWeightManifest { + let cfg = config(); + let mut weight_map = serde_json::Map::new(); + weight_map.insert( + TEXT_EMBEDDING.to_string(), + serde_json::json!("model.safetensors"), + ); + weight_map.insert( + FUSED_MODALITY_EMBEDDING.to_string(), + serde_json::json!("model.safetensors"), + ); + for name in required_body_tensors(&cfg) { + weight_map.insert(name, serde_json::json!("model.safetensors")); + } + HiggsWeightManifest::from_json(&serde_json::json!({ + "metadata": {"total_size": "8489763794"}, + "weight_map": weight_map + })) + .unwrap() + } + + #[test] + fn builds_runtime_load_plan_for_higgs_backbone_and_audio_head() { + let cfg = config(); + let plan = HiggsRuntimeLoadPlan::from_manifest(&cfg, &manifest()).unwrap(); + let summary = plan.summary(); + + assert_eq!(summary.tensors, 399); + assert_eq!(summary.shard_files, 1); + assert_eq!(summary.body_tensors, 397); + assert_eq!(summary.qwen3_backbone_tensors, 398); + assert_eq!(summary.higgs_head_tensors, 1); + + let text = plan.tensor(TEXT_EMBEDDING).unwrap(); + assert_eq!(text.role, TensorRole::TextEmbedding); + assert_eq!(text.loader_slot, "qwen3.embed_tokens"); + assert_eq!(text.shape, [EXPECTED_TEXT_VOCAB_SIZE, HIDDEN_SIZE]); + + let audio = plan.tensor(FUSED_MODALITY_EMBEDDING).unwrap(); + assert_eq!(audio.role, TensorRole::FusedAudioHead); + assert_eq!(audio.loader_slot, "higgs.fused_audio_head"); + assert_eq!(audio.shape, fused_modality_shape()); + + let q_proj = plan + .tensor("body.layers.0.self_attn.q_proj.weight") + .unwrap(); + assert_eq!(q_proj.role, TensorRole::LayerQProj); + assert_eq!(q_proj.loader_slot, "qwen3.layers.0.self_attn.q_proj.weight"); + } + + #[test] + fn load_plan_rejects_missing_required_manifest_tensor() { + let cfg = config(); + let mut manifest = manifest(); + manifest.weight_map.remove(TEXT_EMBEDDING); + let err = HiggsRuntimeLoadPlan::from_manifest(&cfg, &manifest) + .unwrap_err() + .to_string(); + assert!(err.contains(TEXT_EMBEDDING)); + } + + #[test] + fn builds_qwen3_tensor_aliases_without_audio_head() { + let cfg = config(); + let plan = HiggsRuntimeLoadPlan::from_manifest(&cfg, &manifest()).unwrap(); + let aliases = plan.qwen3_tensor_aliases().unwrap(); + + assert_eq!(aliases.len(), 398); + assert_eq!( + aliases.get("model.embed_tokens.weight").unwrap(), + TEXT_EMBEDDING + ); + assert_eq!(aliases.get("model.norm.weight").unwrap(), BODY_NORM); + assert_eq!( + aliases + .get("model.layers.0.self_attn.q_proj.weight") + .unwrap(), + "body.layers.0.self_attn.q_proj.weight" + ); + assert!( + !aliases + .values() + .any(|checkpoint_name| checkpoint_name == FUSED_MODALITY_EMBEDDING) + ); + } +} diff --git a/pegainfer-higgs-audio/src/materialize_qwen3.rs b/pegainfer-higgs-audio/src/materialize_qwen3.rs new file mode 100644 index 000000000..80dca6996 --- /dev/null +++ b/pegainfer-higgs-audio/src/materialize_qwen3.rs @@ -0,0 +1,415 @@ +use std::collections::BTreeMap; +use std::io::Read; +use std::io::Seek; +use std::io::SeekFrom; +use std::io::Write; +use std::path::Path; +use std::path::PathBuf; + +use anyhow::Context; +use anyhow::Result; +use anyhow::ensure; +use serde_json::Value; +use serde_json::json; + +use crate::config::HiggsConfig; +use crate::load_plan::HiggsRuntimeLoadPlan; +use crate::load_plan::PlannedTensor; +use crate::load_plan::qwen3_tensor_name; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MaterializeSummary { + pub output_dir: PathBuf, + pub tensors: usize, + pub payload_bytes: usize, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ConfigViewSummary { + pub output_dir: PathBuf, + pub alias_manifest: PathBuf, + pub aliases: usize, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct SourceTensorHeader { + dtype: String, + shape: Vec, + data_offsets: [usize; 2], +} + +pub fn materialize_qwen3_body_view( + source_model_dir: impl AsRef, + output_dir: impl AsRef, + config: &HiggsConfig, + load_plan: &HiggsRuntimeLoadPlan, +) -> Result { + let source_model_dir = source_model_dir.as_ref(); + let output_dir = output_dir.as_ref(); + std::fs::create_dir_all(output_dir) + .with_context(|| format!("create {}", output_dir.display()))?; + + write_qwen3_config(output_dir, config)?; + write_generation_config(output_dir, config)?; + + let qwen_tensors: Vec<_> = load_plan + .tensors + .iter() + .filter(|tensor| tensor.loader_slot.starts_with("qwen3.")) + .collect(); + ensure!( + qwen_tensors.len() == 398, + "expected 398 Qwen3 backbone tensors, got {}", + qwen_tensors.len() + ); + + let shard_files = qwen_tensors + .iter() + .map(|tensor| tensor.shard_file.as_str()) + .collect::>(); + ensure!( + shard_files.len() == 1, + "Qwen3 body materializer currently expects one Higgs source shard, got {}", + shard_files.len() + ); + let shard_file = shard_files.iter().next().context("missing source shard")?; + let source_safetensors = source_model_dir.join(shard_file); + let output_safetensors = output_dir.join("model.safetensors"); + materialize_safetensors_alias(&source_safetensors, &output_safetensors, &qwen_tensors)?; + + Ok(MaterializeSummary { + output_dir: output_dir.to_path_buf(), + tensors: qwen_tensors.len(), + payload_bytes: qwen_tensors.iter().map(|tensor| tensor.bytes).sum(), + }) +} + +pub fn write_qwen3_config_view( + output_dir: impl AsRef, + config: &HiggsConfig, + load_plan: &HiggsRuntimeLoadPlan, +) -> Result { + let output_dir = output_dir.as_ref(); + std::fs::create_dir_all(output_dir) + .with_context(|| format!("create {}", output_dir.display()))?; + write_qwen3_config(output_dir, config)?; + write_generation_config(output_dir, config)?; + + let aliases = load_plan.qwen3_tensor_aliases()?; + let alias_manifest = output_dir.join("higgs-qwen3-tensor-aliases.json"); + write_json( + &alias_manifest, + &json!({ + "format": "higgs-qwen3-tensor-aliases-v1", + "requested_to_stored": aliases, + }), + )?; + Ok(ConfigViewSummary { + output_dir: output_dir.to_path_buf(), + alias_manifest, + aliases: load_plan.qwen3_tensor_aliases()?.len(), + }) +} + +fn write_qwen3_config(output_dir: &Path, config: &HiggsConfig) -> Result<()> { + let value = json!({ + "hidden_size": config.text.hidden_size, + "intermediate_size": config.text.intermediate_size, + "num_hidden_layers": config.text.num_hidden_layers, + "num_attention_heads": config.text.num_attention_heads, + "num_key_value_heads": config.text.num_key_value_heads, + "head_dim": config.text.head_dim, + "vocab_size": config.text.vocab_size, + "rms_norm_eps": config.text.rms_norm_eps, + "rope_theta": config.text.rope_theta as f32, + "eos_token_id": config.text.eos_token_id, + "tie_word_embeddings": true + }); + write_json(output_dir.join("config.json"), &value) +} + +fn write_generation_config(output_dir: &Path, config: &HiggsConfig) -> Result<()> { + write_json( + output_dir.join("generation_config.json"), + &json!({"eos_token_id": config.text.eos_token_id}), + ) +} + +fn write_json(path: impl AsRef, value: &Value) -> Result<()> { + let bytes = serde_json::to_vec_pretty(value).context("serialize JSON")?; + std::fs::write(path.as_ref(), [bytes, b"\n".to_vec()].concat()) + .with_context(|| format!("write {}", path.as_ref().display())) +} + +fn materialize_safetensors_alias( + source_path: &Path, + output_path: &Path, + tensors: &[&PlannedTensor], +) -> Result<()> { + let mut source = std::fs::File::open(source_path) + .with_context(|| format!("open {}", source_path.display()))?; + let (source_data_start, source_headers) = read_source_headers(&mut source, source_path)?; + let mut aliases = Vec::with_capacity(tensors.len()); + let mut output_offset = 0usize; + for tensor in tensors { + let source_header = source_headers + .get(&tensor.checkpoint_name) + .with_context(|| format!("source header missing {}", tensor.checkpoint_name))?; + ensure!( + source_header.dtype == tensor.dtype, + "{} dtype drift: plan {} source {}", + tensor.checkpoint_name, + tensor.dtype, + source_header.dtype + ); + ensure!( + source_header.shape == tensor.shape, + "{} shape drift: plan {:?} source {:?}", + tensor.checkpoint_name, + tensor.shape, + source_header.shape + ); + let alias = qwen3_tensor_name(&tensor.loader_slot)?; + aliases.push(( + alias, + tensor.checkpoint_name.clone(), + source_header.data_offsets, + output_offset, + output_offset + tensor.bytes, + source_header.dtype.clone(), + source_header.shape.clone(), + )); + output_offset += tensor.bytes; + } + + let mut header = serde_json::Map::new(); + for (alias, _source_name, _source_offsets, start, end, dtype, shape) in &aliases { + header.insert( + alias.clone(), + json!({ + "dtype": dtype, + "shape": shape, + "data_offsets": [start, end], + }), + ); + } + let mut header_bytes = + serde_json::to_vec(&Value::Object(header)).context("serialize header")?; + while !(8 + header_bytes.len()).is_multiple_of(std::mem::align_of::()) { + header_bytes.push(b' '); + } + let mut output = std::fs::File::create(output_path) + .with_context(|| format!("create {}", output_path.display()))?; + output + .write_all(&(header_bytes.len() as u64).to_le_bytes()) + .context("write safetensors header length")?; + output + .write_all(&header_bytes) + .context("write safetensors header")?; + + let mut buffer = vec![0u8; 8 * 1024 * 1024]; + for (_alias, source_name, source_offsets, _start, _end, _dtype, _shape) in &aliases { + let len = source_offsets[1] - source_offsets[0]; + copy_exact_range( + &mut source, + source_data_start + source_offsets[0] as u64, + len, + &mut output, + &mut buffer, + ) + .with_context(|| format!("copy tensor payload for {source_name}"))?; + } + Ok(()) +} + +fn read_source_headers( + source: &mut std::fs::File, + path: &Path, +) -> Result<(u64, BTreeMap)> { + let mut len_bytes = [0u8; 8]; + source + .read_exact(&mut len_bytes) + .with_context(|| format!("read header length from {}", path.display()))?; + let header_len = usize::try_from(u64::from_le_bytes(len_bytes)) + .with_context(|| format!("{} header length does not fit usize", path.display()))?; + let mut header_bytes = vec![0u8; header_len]; + source + .read_exact(&mut header_bytes) + .with_context(|| format!("read header from {}", path.display()))?; + let value: Value = serde_json::from_slice(&header_bytes) + .with_context(|| format!("parse safetensors header from {}", path.display()))?; + let object = value + .as_object() + .with_context(|| format!("{} header is not an object", path.display()))?; + let mut headers = BTreeMap::new(); + for (name, value) in object { + if name == "__metadata__" { + continue; + } + let dtype = value + .get("dtype") + .and_then(Value::as_str) + .with_context(|| format!("{name} missing dtype"))? + .to_string(); + let shape = value + .get("shape") + .and_then(Value::as_array) + .with_context(|| format!("{name} missing shape"))? + .iter() + .map(|dim| { + dim.as_u64() + .context("shape dim not u64") + .and_then(|dim| usize::try_from(dim).context("shape dim does not fit usize")) + }) + .collect::>>()?; + let offsets = value + .get("data_offsets") + .and_then(Value::as_array) + .with_context(|| format!("{name} missing data_offsets"))?; + ensure!(offsets.len() == 2, "{name} data_offsets length must be 2"); + let start = offsets[0] + .as_u64() + .context("start offset not u64") + .and_then(|offset| usize::try_from(offset).context("start does not fit usize"))?; + let end = offsets[1] + .as_u64() + .context("end offset not u64") + .and_then(|offset| usize::try_from(offset).context("end does not fit usize"))?; + ensure!(start <= end, "{name} invalid data_offsets [{start}, {end}]"); + headers.insert( + name.clone(), + SourceTensorHeader { + dtype, + shape, + data_offsets: [start, end], + }, + ); + } + Ok((8 + header_len as u64, headers)) +} + +fn copy_exact_range( + input: &mut std::fs::File, + start: u64, + len: usize, + output: &mut std::fs::File, + buffer: &mut [u8], +) -> Result<()> { + input.seek(SeekFrom::Start(start))?; + let mut remaining = len; + while remaining > 0 { + let chunk = remaining.min(buffer.len()); + input.read_exact(&mut buffer[..chunk])?; + output.write_all(&buffer[..chunk])?; + remaining -= chunk; + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::load_plan::TensorRole; + use crate::weights::BODY_NORM; + use crate::weights::TEXT_EMBEDDING; + + #[test] + fn qwen3_tensor_name_maps_loader_slots() { + assert_eq!( + qwen3_tensor_name("qwen3.embed_tokens").unwrap(), + "model.embed_tokens.weight" + ); + assert_eq!( + qwen3_tensor_name("qwen3.layers.0.self_attn.q_proj.weight").unwrap(), + "model.layers.0.self_attn.q_proj.weight" + ); + } + + #[test] + fn materializes_qwen3_alias_safetensors_from_small_payload() { + let tmp = tempfile::TempDir::new().unwrap(); + let source = tmp.path().join("higgs.safetensors"); + let out = tmp.path().join("qwen3.safetensors"); + let planned = vec![ + planned( + TEXT_EMBEDDING, + TensorRole::TextEmbedding, + "qwen3.embed_tokens", + [2, 2], + ), + planned(BODY_NORM, TensorRole::BodyNorm, "qwen3.norm", [2, 1]), + planned( + "body.layers.0.self_attn.q_proj.weight", + TensorRole::LayerQProj, + "qwen3.layers.0.self_attn.q_proj.weight", + [2, 2], + ), + ]; + write_small_higgs_safetensors(&source, &planned); + let refs: Vec<_> = planned.iter().collect(); + + materialize_safetensors_alias(&source, &out, &refs).unwrap(); + + let bytes = std::fs::read(out).unwrap(); + let header_len = u64::from_le_bytes(bytes[..8].try_into().unwrap()) as usize; + assert_eq!((8 + header_len) % std::mem::align_of::(), 0); + let tensors = safetensors::SafeTensors::deserialize(&bytes).unwrap(); + assert_eq!(tensors.names().len(), 3); + assert_eq!( + tensors.tensor("model.embed_tokens.weight").unwrap().data(), + &[0; 8] + ); + assert_eq!(tensors.tensor("model.norm.weight").unwrap().data(), &[1; 4]); + assert_eq!( + tensors + .tensor("model.layers.0.self_attn.q_proj.weight") + .unwrap() + .data(), + &[2; 8] + ); + assert!(tensors.tensor(TEXT_EMBEDDING).is_err()); + } + + fn planned( + checkpoint_name: &str, + role: TensorRole, + loader_slot: &str, + shape: [usize; 2], + ) -> PlannedTensor { + let elements = shape.iter().product::(); + PlannedTensor { + checkpoint_name: checkpoint_name.to_string(), + shard_file: "model.safetensors".to_string(), + role, + loader_slot: loader_slot.to_string(), + dtype: "BF16", + shape: shape.to_vec(), + elements, + bytes: elements * 2, + } + } + + fn write_small_higgs_safetensors(path: &Path, planned: &[PlannedTensor]) { + let mut header = serde_json::Map::new(); + let mut payload = Vec::new(); + let mut offset = 0usize; + for (idx, tensor) in planned.iter().enumerate() { + header.insert( + tensor.checkpoint_name.clone(), + json!({ + "dtype": tensor.dtype, + "shape": tensor.shape, + "data_offsets": [offset, offset + tensor.bytes] + }), + ); + payload.extend(std::iter::repeat_n(idx as u8, tensor.bytes)); + offset += tensor.bytes; + } + let header = serde_json::to_vec(&Value::Object(header)).unwrap(); + let mut out = Vec::new(); + out.extend_from_slice(&(header.len() as u64).to_le_bytes()); + out.extend_from_slice(&header); + out.extend_from_slice(&payload); + std::fs::write(path, out).unwrap(); + } +} diff --git a/pegainfer-higgs-audio/src/one_step_actual.rs b/pegainfer-higgs-audio/src/one_step_actual.rs new file mode 100644 index 000000000..5b94d30c8 --- /dev/null +++ b/pegainfer-higgs-audio/src/one_step_actual.rs @@ -0,0 +1,629 @@ +use std::borrow::Cow; +use std::collections::BTreeMap; +use std::collections::HashMap; +use std::path::Path; + +use anyhow::Context; +use anyhow::Result; +use anyhow::ensure; +use half::bf16; +use memmap2::Mmap; +use safetensors::Dtype; +use safetensors::SafeTensors; +use safetensors::tensor::TensorView; +use safetensors::tensor::View; + +use crate::compare::AUDIO_ARGMAX_IDS; +use crate::compare::AUDIO_LOGITS_F32; +use crate::compare::AUDIO_TOP64_IDS; +use crate::compare::AUDIO_TOP64_LOGPROBS_F32; +use crate::compare::FINAL_HIDDEN_BF16; +use crate::compare::PROMPT_ATTENTION_MASK; +use crate::compare::PROMPT_INPUT_IDS; +use crate::compare::PROMPT_LENGTHS; +use crate::one_step_golden::CODEBOOK_VOCAB_SIZE; +use crate::one_step_golden::HIDDEN_SIZE; +use crate::one_step_golden::NUM_CODEBOOKS; +use crate::one_step_golden::TOP_K; +use crate::one_step_golden::validate_required_tensors; +use crate::weights::FUSED_MODALITY_EMBEDDING; +use crate::weights::fused_modality_shape; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PromptTensors { + pub input_ids_padded: Vec, + pub attention_mask: Vec, + pub lengths: Vec, +} + +impl PromptTensors { + pub fn prompt_ids(&self) -> Result> { + ensure!( + self.lengths.len() == 1, + "one-step runtime prompt expects exactly one prompt length, got {}", + self.lengths.len() + ); + ensure!( + self.attention_mask.len() == self.input_ids_padded.len(), + "attention mask len {} must match padded ids len {}", + self.attention_mask.len(), + self.input_ids_padded.len() + ); + let len = usize::try_from( + *self + .lengths + .first() + .context("prompt lengths tensor is empty")?, + ) + .context("prompt length must be non-negative")?; + ensure!( + len <= self.input_ids_padded.len(), + "prompt length {len} exceeds padded ids length {}", + self.input_ids_padded.len() + ); + ensure!(len > 0, "prompt length must be positive"); + let mut mask_sum = 0i64; + for (idx, value) in self.attention_mask.iter().enumerate() { + ensure!( + *value == 0 || *value == 1, + "attention mask at index {idx} must be 0/1, got {value}" + ); + mask_sum += *value; + } + ensure!( + mask_sum == len as i64, + "attention mask sum {mask_sum} must match prompt length {len}" + ); + self.input_ids_padded[..len] + .iter() + .map(|value| { + u32::try_from(*value).with_context(|| format!("prompt id {value} is not u32")) + }) + .collect() + } +} + +#[derive(Debug, Clone, PartialEq)] +pub struct OneStepActualSummary { + pub output_path: std::path::PathBuf, + pub prompt_tokens: usize, + pub hidden_values: usize, + pub audio_logits: usize, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct OneStepAudioPrediction { + pub logits: Vec, + pub top_ids: Vec, + pub top_logprobs: Vec, + pub argmax: Vec, +} + +#[derive(Clone)] +pub(crate) struct OwnedTensor { + dtype: Dtype, + shape: Vec, + data: Vec, +} + +impl View for OwnedTensor { + fn dtype(&self) -> Dtype { + self.dtype + } + + fn shape(&self) -> &[usize] { + &self.shape + } + + fn data(&self) -> Cow<'_, [u8]> { + Cow::Borrowed(&self.data) + } + + fn data_len(&self) -> usize { + self.data.len() + } +} + +pub fn load_prompt_from_golden(path: impl AsRef) -> Result { + let path = path.as_ref(); + let bytes = std::fs::read(path).with_context(|| format!("read {}", path.display()))?; + let st = SafeTensors::deserialize(&bytes).context("parse Higgs golden safetensors")?; + validate_required_tensors(&st, "golden")?; + Ok(PromptTensors { + input_ids_padded: i64_values(tensor(&st, PROMPT_INPUT_IDS)?)?, + attention_mask: i64_values(tensor(&st, PROMPT_ATTENTION_MASK)?)?, + lengths: i64_values(tensor(&st, PROMPT_LENGTHS)?)?, + }) +} + +pub fn load_fused_audio_head_bf16(model_dir: impl AsRef) -> Result> { + let model_dir = model_dir.as_ref(); + let index_path = model_dir.join("model.safetensors.index.json"); + let index: serde_json::Value = serde_json::from_slice( + &std::fs::read(&index_path).with_context(|| format!("read {}", index_path.display()))?, + ) + .with_context(|| format!("parse {}", index_path.display()))?; + let shard = index + .get("weight_map") + .and_then(serde_json::Value::as_object) + .and_then(|weight_map| weight_map.get(FUSED_MODALITY_EMBEDDING)) + .and_then(serde_json::Value::as_str) + .with_context(|| format!("index missing {FUSED_MODALITY_EMBEDDING}"))?; + let shard_path = model_dir.join(shard); + let file = std::fs::File::open(&shard_path) + .with_context(|| format!("open {}", shard_path.display()))?; + let mmap = + unsafe { Mmap::map(&file) }.with_context(|| format!("mmap {}", shard_path.display()))?; + let st = SafeTensors::deserialize(&mmap) + .with_context(|| format!("parse {}", shard_path.display()))?; + let tensor = st + .tensor(FUSED_MODALITY_EMBEDDING) + .with_context(|| format!("missing tensor {FUSED_MODALITY_EMBEDDING}"))?; + ensure!( + tensor.dtype() == Dtype::BF16, + "{FUSED_MODALITY_EMBEDDING} must be BF16" + ); + ensure!( + tensor.shape() == fused_modality_shape(), + "{FUSED_MODALITY_EMBEDDING} shape mismatch: expected {:?}, got {:?}", + fused_modality_shape(), + tensor.shape() + ); + bf16_values(tensor) +} + +pub fn write_one_step_actual( + output_path: impl AsRef, + prompt: &PromptTensors, + final_hidden: &[bf16], + audio_head: &[bf16], +) -> Result { + let prediction = compute_one_step_audio_prediction(final_hidden, audio_head)?; + write_one_step_actual_prediction(output_path, prompt, final_hidden, &prediction) +} + +pub fn compute_one_step_audio_prediction( + final_hidden: &[bf16], + audio_head: &[bf16], +) -> Result { + validate_audio_head_inputs(final_hidden, audio_head)?; + let logits = audio_logits_cpu(final_hidden, audio_head); + Ok(OneStepAudioPrediction::from_logits(logits)) +} + +pub fn write_one_step_actual_prediction( + output_path: impl AsRef, + prompt: &PromptTensors, + final_hidden: &[bf16], + prediction: &OneStepAudioPrediction, +) -> Result { + ensure!( + final_hidden.len() == HIDDEN_SIZE, + "final hidden len mismatch: expected {HIDDEN_SIZE}, got {}", + final_hidden.len() + ); + prediction.validate()?; + write_one_step_actual_tensors( + output_path, + prompt, + final_hidden, + &prediction.logits, + &prediction.top_ids, + &prediction.top_logprobs, + &prediction.argmax, + ) +} + +impl OneStepAudioPrediction { + pub fn from_logits(logits: Vec) -> Self { + let (top_ids, top_logprobs, argmax) = audio_topk_and_argmax(&logits); + Self { + logits, + top_ids, + top_logprobs, + argmax, + } + } + + pub fn validate(&self) -> Result<()> { + ensure!( + self.logits.len() == NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE, + "audio logits len mismatch: expected {}, got {}", + NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE, + self.logits.len() + ); + ensure!( + self.top_ids.len() == NUM_CODEBOOKS * TOP_K, + "top id len mismatch: expected {}, got {}", + NUM_CODEBOOKS * TOP_K, + self.top_ids.len() + ); + ensure!( + self.top_logprobs.len() == NUM_CODEBOOKS * TOP_K, + "top logprob len mismatch: expected {}, got {}", + NUM_CODEBOOKS * TOP_K, + self.top_logprobs.len() + ); + ensure!( + self.argmax.len() == NUM_CODEBOOKS, + "argmax len mismatch: expected {NUM_CODEBOOKS}, got {}", + self.argmax.len() + ); + Ok(()) + } +} + +#[cfg(feature = "runtime-qwen3")] +pub fn write_one_step_actual_with_gpu_audio_head( + output_path: impl AsRef, + prompt: &PromptTensors, + final_hidden: &[bf16], + audio_head: &[bf16], + device_ordinal: usize, +) -> Result { + let prediction = + compute_one_step_audio_prediction_gpu_bf16(final_hidden, audio_head, device_ordinal)?; + write_one_step_actual_prediction(output_path, prompt, final_hidden, &prediction) +} + +#[cfg(feature = "runtime-qwen3")] +pub fn compute_one_step_audio_prediction_gpu_bf16( + final_hidden: &[bf16], + audio_head: &[bf16], + device_ordinal: usize, +) -> Result { + validate_audio_head_inputs(final_hidden, audio_head)?; + let logits = audio_logits_gpu_bf16(final_hidden, audio_head, device_ordinal)?; + Ok(OneStepAudioPrediction::from_logits(logits)) +} + +#[cfg(feature = "runtime-qwen3")] +fn audio_logits_gpu_bf16( + final_hidden: &[bf16], + audio_head: &[bf16], + device_ordinal: usize, +) -> Result> { + let ctx = pegainfer_core::tensor::DeviceContext::new_with_device(device_ordinal).with_context( + || format!("create CUDA context for audio head on device {device_ordinal}"), + )?; + let hidden = pegainfer_core::tensor::DeviceVec::from_host(&ctx, final_hidden) + .context("copy final hidden to GPU for Higgs audio head")?; + let head = pegainfer_core::tensor::DeviceMatrix::from_host( + &ctx, + audio_head, + NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE, + HIDDEN_SIZE, + ) + .context("copy fused Higgs audio head to GPU")?; + let logits = pegainfer_core::ops::linear(&ctx, &hidden, &head) + .context("run Higgs fused audio head as CUDA bf16 linear")?; + logits + .to_host(&ctx) + .context("copy Higgs audio logits from GPU") +} + +fn write_one_step_actual_tensors( + output_path: impl AsRef, + prompt: &PromptTensors, + final_hidden: &[bf16], + logits: &[f32], + top_ids: &[i64], + top_logprobs: &[f32], + argmax: &[i64], +) -> Result { + ensure!( + logits.len() == NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE, + "audio logits len mismatch: expected {}, got {}", + NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE, + logits.len() + ); + ensure!( + top_ids.len() == NUM_CODEBOOKS * TOP_K, + "top id len mismatch: expected {}, got {}", + NUM_CODEBOOKS * TOP_K, + top_ids.len() + ); + ensure!( + top_logprobs.len() == NUM_CODEBOOKS * TOP_K, + "top logprob len mismatch: expected {}, got {}", + NUM_CODEBOOKS * TOP_K, + top_logprobs.len() + ); + ensure!( + argmax.len() == NUM_CODEBOOKS, + "argmax len mismatch: expected {NUM_CODEBOOKS}, got {}", + argmax.len() + ); + let output_path = output_path.as_ref(); + let tensors = BTreeMap::from([ + ( + PROMPT_INPUT_IDS.to_string(), + owned_i64( + &[1, prompt.input_ids_padded.len()], + &prompt.input_ids_padded, + ), + ), + ( + PROMPT_ATTENTION_MASK.to_string(), + owned_i64(&[1, prompt.attention_mask.len()], &prompt.attention_mask), + ), + ( + PROMPT_LENGTHS.to_string(), + owned_i64(&[prompt.lengths.len()], &prompt.lengths), + ), + ( + FINAL_HIDDEN_BF16.to_string(), + owned_bf16(&[1, HIDDEN_SIZE], final_hidden), + ), + ( + AUDIO_LOGITS_F32.to_string(), + owned_f32(&[1, NUM_CODEBOOKS, CODEBOOK_VOCAB_SIZE], logits), + ), + ( + AUDIO_TOP64_IDS.to_string(), + owned_i64(&[1, NUM_CODEBOOKS, TOP_K], top_ids), + ), + ( + AUDIO_TOP64_LOGPROBS_F32.to_string(), + owned_f32(&[1, NUM_CODEBOOKS, TOP_K], top_logprobs), + ), + ( + AUDIO_ARGMAX_IDS.to_string(), + owned_i64(&[1, NUM_CODEBOOKS], argmax), + ), + ]); + let metadata = HashMap::from([( + "fixture_kind".to_string(), + "higgs-one-step-audio-logits-actual".to_string(), + )]); + safetensors::serialize_to_file(tensors, Some(metadata), output_path) + .with_context(|| format!("write {}", output_path.display()))?; + + Ok(OneStepActualSummary { + output_path: output_path.to_path_buf(), + prompt_tokens: prompt.prompt_ids()?.len(), + hidden_values: final_hidden.len(), + audio_logits: logits.len(), + }) +} + +fn validate_audio_head_inputs(final_hidden: &[bf16], audio_head: &[bf16]) -> Result<()> { + ensure!( + final_hidden.len() == HIDDEN_SIZE, + "final hidden len mismatch: expected {HIDDEN_SIZE}, got {}", + final_hidden.len() + ); + ensure!( + audio_head.len() == NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE * HIDDEN_SIZE, + "audio head len mismatch: expected {}, got {}", + NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE * HIDDEN_SIZE, + audio_head.len() + ); + Ok(()) +} + +fn audio_logits_cpu(final_hidden: &[bf16], audio_head: &[bf16]) -> Vec { + let hidden: Vec = final_hidden.iter().map(|value| value.to_f32()).collect(); + audio_head + .chunks_exact(HIDDEN_SIZE) + .map(|row| { + row.iter() + .zip(&hidden) + .map(|(weight, hidden)| weight.to_f32() * hidden) + .sum() + }) + .collect() +} + +fn audio_topk_and_argmax(logits: &[f32]) -> (Vec, Vec, Vec) { + let mut top_ids = Vec::with_capacity(NUM_CODEBOOKS * TOP_K); + let mut top_logprobs = Vec::with_capacity(NUM_CODEBOOKS * TOP_K); + let mut argmax = Vec::with_capacity(NUM_CODEBOOKS); + for row in logits.chunks_exact(CODEBOOK_VOCAB_SIZE) { + let max = row.iter().copied().fold(f32::NEG_INFINITY, f32::max); + let logsumexp = max + + row + .iter() + .map(|value| (*value - max).exp()) + .sum::() + .ln(); + let mut indexed: Vec<_> = row.iter().copied().enumerate().collect(); + indexed.sort_by(|left, right| { + right + .1 + .total_cmp(&left.1) + .then_with(|| left.0.cmp(&right.0)) + }); + argmax.push(indexed[0].0 as i64); + for (idx, value) in indexed.into_iter().take(TOP_K) { + top_ids.push(idx as i64); + top_logprobs.push(value - logsumexp); + } + } + (top_ids, top_logprobs, argmax) +} + +pub(crate) fn owned_i64(shape: &[usize], values: &[i64]) -> OwnedTensor { + OwnedTensor { + dtype: Dtype::I64, + shape: shape.to_vec(), + data: values + .iter() + .flat_map(|value| value.to_le_bytes()) + .collect(), + } +} + +fn owned_f32(shape: &[usize], values: &[f32]) -> OwnedTensor { + OwnedTensor { + dtype: Dtype::F32, + shape: shape.to_vec(), + data: values + .iter() + .flat_map(|value| value.to_le_bytes()) + .collect(), + } +} + +pub(crate) fn owned_bf16(shape: &[usize], values: &[bf16]) -> OwnedTensor { + OwnedTensor { + dtype: Dtype::BF16, + shape: shape.to_vec(), + data: values + .iter() + .flat_map(|value| value.to_bits().to_le_bytes()) + .collect(), + } +} + +fn tensor<'a>(st: &'a SafeTensors, name: &str) -> Result> { + st.tensor(name) + .with_context(|| format!("safetensors missing tensor {name}")) +} + +fn i64_values(tensor: TensorView<'_>) -> Result> { + ensure!(tensor.dtype() == Dtype::I64, "tensor must be I64"); + ensure!( + tensor.data().len().is_multiple_of(8), + "I64 tensor byte length {} is not divisible by 8", + tensor.data().len() + ); + Ok(tensor + .data() + .chunks_exact(8) + .map(|bytes| { + i64::from_le_bytes([ + bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7], + ]) + }) + .collect()) +} + +fn bf16_values(tensor: TensorView<'_>) -> Result> { + ensure!(tensor.dtype() == Dtype::BF16, "tensor must be BF16"); + ensure!( + tensor.data().len().is_multiple_of(2), + "BF16 tensor byte length {} is not divisible by 2", + tensor.data().len() + ); + Ok(tensor + .data() + .chunks_exact(2) + .map(|bytes| bf16::from_bits(u16::from_le_bytes([bytes[0], bytes[1]]))) + .collect()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::compare::OneStepTolerances; + use crate::compare::compare_one_step_files; + + #[test] + fn prompt_ids_respect_recorded_length() { + let prompt = PromptTensors { + input_ids_padded: vec![10, 20, 30, 0], + attention_mask: vec![1, 1, 1, 0], + lengths: vec![3], + }; + assert_eq!(prompt.prompt_ids().unwrap(), vec![10, 20, 30]); + } + + #[test] + fn prompt_ids_reject_multi_prompt_surface() { + let prompt = PromptTensors { + input_ids_padded: vec![10, 20, 30, 40], + attention_mask: vec![1, 1, 1, 1], + lengths: vec![2, 2], + }; + let err = prompt.prompt_ids().unwrap_err().to_string(); + assert!(err.contains("expects exactly one prompt length")); + } + + #[test] + fn prompt_ids_reject_attention_mask_sum_mismatch() { + let prompt = PromptTensors { + input_ids_padded: vec![10, 20, 30, 0], + attention_mask: vec![1, 1, 0, 0], + lengths: vec![3], + }; + let err = prompt.prompt_ids().unwrap_err().to_string(); + assert!(err.contains("attention mask sum 2 must match prompt length 3")); + } + + #[test] + fn prompt_ids_reject_non_binary_attention_mask() { + let prompt = PromptTensors { + input_ids_padded: vec![10, 20, 30, 0], + attention_mask: vec![1, 2, 0, 0], + lengths: vec![3], + }; + let err = prompt.prompt_ids().unwrap_err().to_string(); + assert!(err.contains("attention mask at index 1 must be 0/1")); + } + + #[test] + fn actual_writer_emits_comparator_schema() { + let tmp = tempfile::NamedTempFile::new().unwrap(); + let prompt = PromptTensors { + input_ids_padded: vec![1; 10], + attention_mask: vec![1; 10], + lengths: vec![10], + }; + let hidden = vec![bf16::from_f32(0.0); HIDDEN_SIZE]; + let mut head = vec![bf16::from_f32(0.0); NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE * HIDDEN_SIZE]; + for codebook in 0..NUM_CODEBOOKS { + let row = codebook * CODEBOOK_VOCAB_SIZE; + head[row * HIDDEN_SIZE] = bf16::from_f32(1.0); + } + write_one_step_actual(tmp.path(), &prompt, &hidden, &head).unwrap(); + let err = compare_one_step_files(tmp.path(), tmp.path(), OneStepTolerances::default()) + .err() + .map(|err| err.to_string()); + assert_eq!(err, None); + } + + #[test] + fn audio_prediction_is_reusable_before_writing() { + let hidden = vec![bf16::from_f32(1.0); HIDDEN_SIZE]; + let mut head = vec![bf16::from_f32(0.0); NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE * HIDDEN_SIZE]; + for codebook in 0..NUM_CODEBOOKS { + let row = codebook * CODEBOOK_VOCAB_SIZE + codebook; + head[row * HIDDEN_SIZE] = bf16::from_f32(1.0); + } + + let prediction = compute_one_step_audio_prediction(&hidden, &head).unwrap(); + + prediction.validate().unwrap(); + assert_eq!(prediction.logits.len(), NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE); + assert_eq!( + prediction.argmax, + (0..NUM_CODEBOOKS as i64).collect::>() + ); + } + + #[test] + fn audio_prediction_validates_shape_contract() { + let prediction = OneStepAudioPrediction { + logits: vec![0.0; NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE - 1], + top_ids: vec![0; NUM_CODEBOOKS * TOP_K], + top_logprobs: vec![0.0; NUM_CODEBOOKS * TOP_K], + argmax: vec![0; NUM_CODEBOOKS], + }; + + let error = prediction.validate().unwrap_err().to_string(); + assert!(error.contains("audio logits len mismatch")); + } + + #[test] + fn topk_uses_local_audio_token_ids() { + let mut logits = vec![0.0; NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE]; + logits[5] = 10.0; + logits[CODEBOOK_VOCAB_SIZE + 7] = 11.0; + let (top_ids, _top_logprobs, argmax) = audio_topk_and_argmax(&logits); + assert_eq!(argmax[0], 5); + assert_eq!(argmax[1], 7); + assert_eq!(top_ids[0], 5); + assert_eq!(top_ids[TOP_K], 7); + } +} diff --git a/pegainfer-higgs-audio/src/one_step_golden.rs b/pegainfer-higgs-audio/src/one_step_golden.rs new file mode 100644 index 000000000..1ada5c6a4 --- /dev/null +++ b/pegainfer-higgs-audio/src/one_step_golden.rs @@ -0,0 +1,214 @@ +use std::collections::HashMap; +use std::path::Path; + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use safetensors::Dtype; +use safetensors::SafeTensors; +use sha2::Digest; +use sha2::Sha256; + +pub const FIXTURE_KIND: &str = "higgs-one-step-audio-logits-golden"; +pub const MODEL_ID: &str = "bosonai/higgs-tts-3-4b"; +pub const MODEL_REVISION: &str = "7556c17e05201fccd9c8cc120bc216dcc7b5d561"; +pub const NUM_CODEBOOKS: usize = 8; +pub const CODEBOOK_VOCAB_SIZE: usize = 1026; +pub const HIDDEN_SIZE: usize = 2560; +pub const TOP_K: usize = 64; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct TensorSpec { + pub name: &'static str, + pub dtype: Dtype, + pub shape: &'static [usize], +} + +pub const REQUIRED_TENSORS: &[TensorSpec] = &[ + TensorSpec { + name: "prompt.input_ids_padded", + dtype: Dtype::I64, + shape: &[1, 10], + }, + TensorSpec { + name: "prompt.attention_mask", + dtype: Dtype::I64, + shape: &[1, 10], + }, + TensorSpec { + name: "prompt.lengths", + dtype: Dtype::I64, + shape: &[1], + }, + TensorSpec { + name: "final_hidden.bf16", + dtype: Dtype::BF16, + shape: &[1, HIDDEN_SIZE], + }, + TensorSpec { + name: "audio_logits.f32", + dtype: Dtype::F32, + shape: &[1, NUM_CODEBOOKS, CODEBOOK_VOCAB_SIZE], + }, + TensorSpec { + name: "audio_top64.ids", + dtype: Dtype::I64, + shape: &[1, NUM_CODEBOOKS, TOP_K], + }, + TensorSpec { + name: "audio_top64.logprobs.f32", + dtype: Dtype::F32, + shape: &[1, NUM_CODEBOOKS, TOP_K], + }, + TensorSpec { + name: "audio_argmax.ids", + dtype: Dtype::I64, + shape: &[1, NUM_CODEBOOKS], + }, +]; + +#[derive(Debug, Clone)] +pub struct GoldenContract { + pub metadata: HashMap, + pub sha256: String, + pub bytes: usize, +} + +pub fn load_and_validate(path: impl AsRef) -> Result { + let path = path.as_ref(); + let bytes = std::fs::read(path).with_context(|| format!("read {}", path.display()))?; + let sha256 = sha256_hex(&bytes); + let metadata = safetensors_metadata(&bytes)?; + validate_metadata(&metadata)?; + let st = SafeTensors::deserialize(&bytes).context("parse Higgs golden safetensors")?; + validate_required_tensors(&st, "golden")?; + Ok(GoldenContract { + metadata, + sha256, + bytes: bytes.len(), + }) +} + +fn validate_metadata(metadata: &HashMap) -> Result<()> { + require_metadata(metadata, "fixture_kind", FIXTURE_KIND)?; + require_metadata(metadata, "model_id", MODEL_ID)?; + require_metadata(metadata, "model_revision", MODEL_REVISION)?; + require_metadata(metadata, "schema_version", "1")?; + require_metadata(metadata, "num_codebooks", &NUM_CODEBOOKS.to_string())?; + require_metadata( + metadata, + "codebook_vocab_size", + &CODEBOOK_VOCAB_SIZE.to_string(), + )?; + require_metadata(metadata, "hidden_size", &HIDDEN_SIZE.to_string())?; + let reference = metadata + .get("reference") + .context("golden metadata missing reference")?; + if !reference.contains("SGLang-Omni Higgs prompt builder") { + bail!("golden reference must record SGLang-Omni Higgs prompt semantics"); + } + let files = metadata + .get("sglang_omni_reference_files") + .context("golden metadata missing sglang_omni_reference_files")?; + for expected in [ + "sglang_omni/models/higgs_tts/text_tokenizer.py", + "sglang_omni/models/higgs_tts/modeling.py", + "sglang_omni/models/higgs_tts/model.py", + ] { + if !files.contains(expected) { + bail!("golden metadata reference files missing {expected}"); + } + } + Ok(()) +} + +pub fn validate_required_tensors(st: &SafeTensors, label: &str) -> Result<()> { + for spec in REQUIRED_TENSORS { + let tensor = st + .tensor(spec.name) + .with_context(|| format!("{label} missing tensor {}", spec.name))?; + if tensor.dtype() != spec.dtype { + bail!( + "{label} tensor {} dtype mismatch: expected {:?}, got {:?}", + spec.name, + spec.dtype, + tensor.dtype() + ); + } + if tensor.shape() != spec.shape { + bail!( + "{label} tensor {} shape mismatch: expected {:?}, got {:?}", + spec.name, + spec.shape, + tensor.shape() + ); + } + } + Ok(()) +} + +fn require_metadata(metadata: &HashMap, key: &str, expected: &str) -> Result<()> { + let actual = metadata + .get(key) + .with_context(|| format!("golden metadata missing {key}"))?; + if actual != expected { + bail!("golden metadata {key} mismatch: expected {expected}, got {actual}"); + } + Ok(()) +} + +fn safetensors_metadata(bytes: &[u8]) -> Result> { + let header_len_bytes: [u8; 8] = bytes + .get(..8) + .context("safetensors file missing 8-byte header length")? + .try_into() + .expect("slice length checked"); + let header_len = u64::from_le_bytes(header_len_bytes) as usize; + let header = bytes + .get(8..8 + header_len) + .context("safetensors file missing JSON header")?; + let value: serde_json::Value = + serde_json::from_slice(header).context("parse safetensors JSON header")?; + Ok(value + .get("__metadata__") + .and_then(serde_json::Value::as_object) + .map(|metadata| { + metadata + .iter() + .filter_map(|(key, value)| { + value.as_str().map(|value| (key.clone(), value.to_string())) + }) + .collect() + }) + .unwrap_or_default()) +} + +fn sha256_hex(bytes: &[u8]) -> String { + let mut digest = Sha256::new(); + digest.update(bytes); + digest + .finalize() + .iter() + .map(|byte| format!("{byte:02x}")) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + + const GOLDEN: &str = concat!( + env!("CARGO_MANIFEST_DIR"), + "/../test_data/higgs-one-step-audio-logits.safetensors" + ); + + #[test] + fn committed_higgs_one_step_golden_has_expected_contract() { + let contract = load_and_validate(GOLDEN).expect("validate committed Higgs golden"); + assert_eq!( + contract.sha256, + "bbaae8018759b7e8f26d2acfb1aefb5bee3e5099d47573bb4ce3c980b6096684" + ); + assert_eq!(contract.bytes, 46_064); + } +} diff --git a/pegainfer-higgs-audio/src/runtime_bridge.rs b/pegainfer-higgs-audio/src/runtime_bridge.rs new file mode 100644 index 000000000..a11a38cb5 --- /dev/null +++ b/pegainfer-higgs-audio/src/runtime_bridge.rs @@ -0,0 +1,258 @@ +use std::path::Path; + +use anyhow::Context; +use anyhow::Result; +use half::bf16; +use pegainfer_core::weight_loader::TensorNameAliases; +use pegainfer_qwen3::runtime::Qwen3Executor; +use pegainfer_qwen3::runtime::RequestId; + +use crate::config::HiggsConfig; +use crate::load_plan::HiggsRuntimeLoadPlan; +use crate::materialize_qwen3::write_qwen3_config_view; +use crate::one_step_actual::OneStepActualSummary; +use crate::one_step_actual::OneStepAudioPrediction; +use crate::one_step_actual::PromptTensors; +use crate::one_step_actual::compute_one_step_audio_prediction; +use crate::one_step_actual::compute_one_step_audio_prediction_gpu_bf16; +use crate::one_step_actual::load_fused_audio_head_bf16; +use crate::one_step_actual::load_prompt_from_golden; +use crate::one_step_actual::write_one_step_actual_prediction; +use crate::weights::HiggsWeightManifest; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum AudioHeadBackend { + CudaBf16, + CpuFp32, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum HiggsRuntimeSource<'a> { + Qwen3BodyView { qwen3_body_dir: &'a Path }, + Qwen3ConfigAlias { qwen3_config_dir: &'a Path }, + AutoConfigAlias { qwen3_config_dir: &'a Path }, +} + +/// Higgs Audio runtime surface backed by the existing Qwen3 executor. +/// +/// The current implementation owns prefill and prompt-session smoke paths. Full +/// audio decode continuation is intentionally not exposed until the Higgs crate +/// owns the audio-codebook feedback semantics. +pub struct HiggsAudioRuntime { + executor: Qwen3Executor, + audio_head: Vec, + audio_head_backend: AudioHeadBackend, + device_ordinal: usize, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct HiggsAudioPrefill { + pub prompt_tokens: usize, + pub final_hidden_bf16: Vec, + pub audio: OneStepAudioPrediction, +} + +/// Compatibility alias for early one-step gate callers. +pub type HiggsOneStepRuntime = HiggsAudioRuntime; + +/// Compatibility alias for early one-step gate callers. +pub type HiggsOneStepPrefill = HiggsAudioPrefill; + +/// Higgs-owned handle for a retained prompt KV session. +/// +/// The backing executor currently stores the session under a Qwen3 request id, +/// but callers should treat this as a Higgs Audio session id. Audio-codebook +/// continuation is intentionally not exposed through this handle yet. +#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)] +pub struct HiggsPromptSession { + request_id: RequestId, +} + +impl HiggsPromptSession { + pub fn new(id: u64) -> Self { + Self { + request_id: RequestId::new(id), + } + } + + pub fn id(self) -> u64 { + self.request_id.get() + } + + fn request_id(self) -> RequestId { + self.request_id + } +} + +impl From for HiggsPromptSession { + fn from(request_id: RequestId) -> Self { + Self { request_id } + } +} + +#[derive(Debug, Clone, PartialEq)] +pub struct HiggsPromptSessionPrefill { + pub session: HiggsPromptSession, + pub prompt_tokens: usize, + pub final_hidden_bf16: Vec, + pub audio: OneStepAudioPrediction, +} + +impl HiggsAudioRuntime { + pub fn from_model_dir( + model_dir: impl AsRef, + source: HiggsRuntimeSource<'_>, + audio_head_backend: AudioHeadBackend, + device_ordinal: usize, + ) -> Result { + let model_dir = model_dir.as_ref(); + let executor = load_qwen3_executor(model_dir, source, device_ordinal)?; + let audio_head = load_fused_audio_head_bf16(model_dir)?; + Ok(Self { + executor, + audio_head, + audio_head_backend, + device_ordinal, + }) + } + + pub fn dump_one_step_actual( + &mut self, + golden: impl AsRef, + out: impl AsRef, + ) -> Result { + let prompt = load_prompt_from_golden(golden)?; + let prompt_ids = prompt.prompt_ids()?; + let prefill = self.prefill_audio_from_prompt_ids(&prompt_ids)?; + write_one_step_actual_prediction(out, &prompt, &prefill.final_hidden_bf16, &prefill.audio) + } + + pub fn prefill_audio_from_prompt( + &mut self, + prompt: &PromptTensors, + ) -> Result { + let prompt_ids = prompt.prompt_ids()?; + self.prefill_audio_from_prompt_ids(&prompt_ids) + } + + pub fn prefill_audio_from_prompt_ids( + &mut self, + prompt_ids: &[u32], + ) -> Result { + let hidden = self + .executor + .prefill_last_hidden_bf16(prompt_ids.to_vec())? + .hidden_bf16; + let audio = match self.audio_head_backend { + AudioHeadBackend::CudaBf16 => compute_one_step_audio_prediction_gpu_bf16( + &hidden, + &self.audio_head, + self.device_ordinal, + )?, + AudioHeadBackend::CpuFp32 => { + compute_one_step_audio_prediction(&hidden, &self.audio_head)? + } + }; + Ok(HiggsAudioPrefill { + prompt_tokens: prompt_ids.len(), + final_hidden_bf16: hidden, + audio, + }) + } + + pub fn prefill_prompt_session_from_prompt_ids( + &mut self, + request_id: RequestId, + prompt_ids: &[u32], + ) -> Result { + self.prefill_prompt_session(request_id.into(), prompt_ids) + } + + pub fn prefill_prompt_session( + &mut self, + session: HiggsPromptSession, + prompt_ids: &[u32], + ) -> Result { + let retained = self + .executor + .prefill_last_hidden_bf16_retained_prompt(session.request_id(), prompt_ids.to_vec())?; + let audio = match self.audio_head_backend { + AudioHeadBackend::CudaBf16 => compute_one_step_audio_prediction_gpu_bf16( + &retained.hidden_bf16, + &self.audio_head, + self.device_ordinal, + )?, + AudioHeadBackend::CpuFp32 => { + compute_one_step_audio_prediction(&retained.hidden_bf16, &self.audio_head)? + } + }; + Ok(HiggsPromptSessionPrefill { + session: retained.request_id.into(), + prompt_tokens: prompt_ids.len(), + final_hidden_bf16: retained.hidden_bf16, + audio, + }) + } + + pub fn drop_prompt_session(&mut self, session: impl Into) -> Result<()> { + self.executor.drop_request(session.into().request_id()) + } +} + +fn load_qwen3_executor( + model_dir: &Path, + source: HiggsRuntimeSource<'_>, + device_ordinal: usize, +) -> Result { + match source { + HiggsRuntimeSource::Qwen3BodyView { qwen3_body_dir } => { + let qwen3_body_dir = path_str(qwen3_body_dir, "qwen3 body dir")?; + Qwen3Executor::from_runtime(qwen3_body_dir, false, &[device_ordinal]) + } + HiggsRuntimeSource::Qwen3ConfigAlias { qwen3_config_dir } => { + load_qwen3_executor_from_alias_config(model_dir, qwen3_config_dir, device_ordinal) + } + HiggsRuntimeSource::AutoConfigAlias { qwen3_config_dir } => { + prepare_qwen3_config_view(model_dir, qwen3_config_dir)?; + load_qwen3_executor_from_alias_config(model_dir, qwen3_config_dir, device_ordinal) + } + } +} + +fn load_qwen3_executor_from_alias_config( + model_dir: &Path, + qwen3_config_dir: &Path, + device_ordinal: usize, +) -> Result { + let qwen3_config_dir = path_str(qwen3_config_dir, "qwen3 config dir")?; + let model_dir_str = path_str(model_dir, "model dir")?; + Qwen3Executor::from_runtime_with_weight_source( + qwen3_config_dir, + Some(model_dir_str), + qwen3_tensor_name_aliases(model_dir)?, + false, + &[device_ordinal], + ) +} + +fn prepare_qwen3_config_view(model_dir: &Path, qwen3_config_dir: &Path) -> Result<()> { + let config = HiggsConfig::from_model_dir(model_dir)?; + let manifest = HiggsWeightManifest::from_model_dir(model_dir)?; + let plan = HiggsRuntimeLoadPlan::from_manifest(&config, &manifest)?; + write_qwen3_config_view(qwen3_config_dir, &config, &plan)?; + Ok(()) +} + +fn qwen3_tensor_name_aliases(model_dir: &Path) -> Result { + let config = HiggsConfig::from_model_dir(model_dir)?; + let manifest = HiggsWeightManifest::from_model_dir(model_dir)?; + let plan = HiggsRuntimeLoadPlan::from_manifest(&config, &manifest)?; + Ok(TensorNameAliases::new( + plan.qwen3_tensor_aliases()?.into_iter().collect(), + )) +} + +fn path_str<'a>(path: &'a Path, label: &str) -> Result<&'a str> { + path.to_str() + .with_context(|| format!("{label} must be valid UTF-8")) +} diff --git a/pegainfer-higgs-audio/src/runtime_source.rs b/pegainfer-higgs-audio/src/runtime_source.rs new file mode 100644 index 000000000..7ffb64fc3 --- /dev/null +++ b/pegainfer-higgs-audio/src/runtime_source.rs @@ -0,0 +1,100 @@ +use std::path::Path; +use std::path::PathBuf; + +use anyhow::Result; +use anyhow::bail; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum Qwen3RuntimeSourcePath<'a> { + BodyView(&'a Path), + ConfigAlias(&'a Path), + AutoConfigAlias(PathBuf), +} + +pub fn select_qwen3_runtime_source<'a>( + qwen3_body_dir: Option<&'a Path>, + qwen3_config_dir: Option<&'a Path>, + out: &Path, +) -> Result> { + match (qwen3_body_dir, qwen3_config_dir) { + (Some(qwen3_body_dir), None) => Ok(Qwen3RuntimeSourcePath::BodyView(qwen3_body_dir)), + (None, Some(qwen3_config_dir)) => Ok(Qwen3RuntimeSourcePath::ConfigAlias(qwen3_config_dir)), + (None, None) => Ok(Qwen3RuntimeSourcePath::AutoConfigAlias( + default_qwen3_config_dir(out), + )), + (Some(_), Some(_)) => bail!("choose only one Qwen3 runtime source"), + } +} + +pub fn default_qwen3_config_dir(out: &Path) -> PathBuf { + out.parent() + .map(|parent| parent.join("higgs-qwen3-config-view")) + .unwrap_or_else(|| PathBuf::from("higgs-qwen3-config-view")) +} + +#[cfg(test)] +mod tests { + use std::path::Path; + + use super::Qwen3RuntimeSourcePath; + use super::default_qwen3_config_dir; + use super::select_qwen3_runtime_source; + + #[test] + fn default_config_view_lives_next_to_actual_output() { + assert_eq!( + default_qwen3_config_dir(Path::new("/tmp/higgs/actual/out.safetensors")), + Path::new("/tmp/higgs/actual/higgs-qwen3-config-view") + ); + } + + #[test] + fn default_config_view_falls_back_for_bare_output_name() { + assert_eq!( + default_qwen3_config_dir(Path::new("out.safetensors")), + Path::new("higgs-qwen3-config-view") + ); + } + + #[test] + fn source_selection_prefers_explicit_body_view() { + let body = Path::new("/tmp/qwen3-body-view"); + let selected = + select_qwen3_runtime_source(Some(body), None, Path::new("/tmp/out.safetensors")) + .unwrap(); + assert_eq!(selected, Qwen3RuntimeSourcePath::BodyView(body)); + } + + #[test] + fn source_selection_prefers_explicit_config_alias() { + let config = Path::new("/tmp/qwen3-config-view"); + let selected = + select_qwen3_runtime_source(None, Some(config), Path::new("/tmp/out.safetensors")) + .unwrap(); + assert_eq!(selected, Qwen3RuntimeSourcePath::ConfigAlias(config)); + } + + #[test] + fn source_selection_defaults_to_auto_config_alias() { + let selected = + select_qwen3_runtime_source(None, None, Path::new("/tmp/higgs/actual/out.safetensors")) + .unwrap(); + assert_eq!( + selected, + Qwen3RuntimeSourcePath::AutoConfigAlias( + Path::new("/tmp/higgs/actual/higgs-qwen3-config-view").to_path_buf() + ) + ); + } + + #[test] + fn source_selection_rejects_ambiguous_runtime_source() { + let error = select_qwen3_runtime_source( + Some(Path::new("/tmp/body")), + Some(Path::new("/tmp/config")), + Path::new("/tmp/out.safetensors"), + ) + .unwrap_err(); + assert!(error.to_string().contains("choose only one")); + } +} diff --git a/pegainfer-higgs-audio/src/weights.rs b/pegainfer-higgs-audio/src/weights.rs new file mode 100644 index 000000000..5b51f7385 --- /dev/null +++ b/pegainfer-higgs-audio/src/weights.rs @@ -0,0 +1,531 @@ +use std::collections::BTreeMap; +use std::collections::BTreeSet; +use std::collections::HashMap; +use std::io::Read; +use std::path::Path; + +use anyhow::Context; +use anyhow::Result; +use anyhow::bail; +use anyhow::ensure; +use serde_json::Value; + +use crate::config::HiggsConfig; +use crate::one_step_golden::CODEBOOK_VOCAB_SIZE; +use crate::one_step_golden::HIDDEN_SIZE; +use crate::one_step_golden::NUM_CODEBOOKS; + +pub const TEXT_EMBEDDING: &str = "tied.embedding.text_embedding.weight"; +pub const FUSED_MODALITY_EMBEDDING: &str = "tied.embedding.modality_embeddings.0.embedding.weight"; +pub const BODY_NORM: &str = "body.norm.weight"; + +#[derive(Debug, Clone)] +pub struct HiggsWeightManifest { + pub total_size: Option, + pub weight_map: HashMap, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ManifestSummary { + pub total_tensors: usize, + pub body_tensors: usize, + pub decoder_only_tensors: usize, + pub has_text_embedding: bool, + pub has_fused_modality_embedding: bool, + pub has_separate_audio_head: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct TensorHeaderSpec { + pub name: String, + pub dtype: &'static str, + pub shape: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct CheckpointHeaderSummary { + pub files_checked: usize, + pub tensors_checked: usize, + pub bf16_tensors_checked: usize, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct TensorHeader { + dtype: String, + shape: Vec, + byte_range: [usize; 2], +} + +impl HiggsWeightManifest { + pub fn from_model_dir(model_dir: impl AsRef) -> Result { + let path = model_dir.as_ref().join("model.safetensors.index.json"); + let value: Value = serde_json::from_slice( + &std::fs::read(&path).with_context(|| format!("read {}", path.display()))?, + ) + .with_context(|| format!("parse {}", path.display()))?; + Self::from_json(&value) + } + + pub fn from_json(value: &Value) -> Result { + let total_size = value + .get("metadata") + .and_then(|metadata| metadata.get("total_size")) + .and_then(|total| match total { + Value::Number(n) => n.as_u64(), + Value::String(s) => s.parse().ok(), + _ => None, + }); + let raw = value + .get("weight_map") + .and_then(Value::as_object) + .context("index missing weight_map")?; + let mut weight_map = HashMap::with_capacity(raw.len()); + for (key, value) in raw { + let file = value + .as_str() + .with_context(|| format!("weight_map entry {key} is not a string"))?; + weight_map.insert(key.clone(), file.to_string()); + } + Ok(Self { + total_size, + weight_map, + }) + } + + pub fn summary(&self) -> ManifestSummary { + ManifestSummary { + total_tensors: self.weight_map.len(), + body_tensors: self + .weight_map + .keys() + .filter(|name| name.starts_with("body.")) + .count(), + decoder_only_tensors: self + .weight_map + .keys() + .filter(|name| is_decoder_only_tensor(name)) + .count(), + has_text_embedding: self.weight_map.contains_key(TEXT_EMBEDDING), + has_fused_modality_embedding: self.weight_map.contains_key(FUSED_MODALITY_EMBEDDING), + has_separate_audio_head: self.weight_map.keys().any(|name| { + name.starts_with("tied.head.modality") || name.starts_with("tied.head.audio") + }), + } + } + + pub fn validate_for_config(&self, config: &HiggsConfig) -> Result { + let summary = self.summary(); + if !summary.has_text_embedding { + bail!("Higgs manifest missing {TEXT_EMBEDDING}"); + } + if !summary.has_fused_modality_embedding { + bail!("Higgs manifest missing {FUSED_MODALITY_EMBEDDING}"); + } + if summary.has_separate_audio_head { + bail!( + "Higgs manifest unexpectedly contains a separate audio head; current checkpoint ties the fused modality head" + ); + } + let required = required_body_tensors(config); + let missing: Vec<_> = required + .iter() + .filter(|name| !self.weight_map.contains_key(*name)) + .cloned() + .collect(); + if !missing.is_empty() { + bail!( + "Higgs manifest missing {} required body tensor(s): {:?}", + missing.len(), + &missing[..missing.len().min(8)] + ); + } + if summary.body_tensors != required.len() { + bail!( + "Higgs manifest body tensor count mismatch: expected {}, got {}", + required.len(), + summary.body_tensors + ); + } + Ok(summary) + } +} + +pub fn validate_checkpoint_headers( + model_dir: impl AsRef, + config: &HiggsConfig, + manifest: &HiggsWeightManifest, +) -> Result { + let expected = expected_checkpoint_tensors(config); + let mut by_file: BTreeMap> = BTreeMap::new(); + for spec in &expected { + let file = manifest + .weight_map + .get(&spec.name) + .with_context(|| format!("manifest missing expected tensor {}", spec.name))?; + by_file.entry(file.clone()).or_default().push(spec); + } + + let mut tensors_checked = 0usize; + let mut bf16_tensors_checked = 0usize; + for (file, specs) in &by_file { + let path = model_dir.as_ref().join(file); + let header = read_safetensors_header(&path)?; + for spec in specs { + let actual = header + .get(&spec.name) + .with_context(|| format!("{} missing tensor {}", path.display(), spec.name))?; + ensure!( + actual.dtype == spec.dtype, + "{} dtype mismatch: expected {}, got {}", + spec.name, + spec.dtype, + actual.dtype + ); + ensure!( + actual.shape == spec.shape, + "{} shape mismatch: expected {:?}, got {:?}", + spec.name, + spec.shape, + actual.shape + ); + ensure!( + actual.byte_range[0] <= actual.byte_range[1], + "{} has invalid data_offsets {:?}", + spec.name, + actual.byte_range + ); + tensors_checked += 1; + if actual.dtype == "BF16" { + bf16_tensors_checked += 1; + } + } + } + + Ok(CheckpointHeaderSummary { + files_checked: by_file.len(), + tensors_checked, + bf16_tensors_checked, + }) +} + +pub fn expected_checkpoint_tensors(config: &HiggsConfig) -> Vec { + let hidden = config.text.hidden_size; + let head_dim = config.text.head_dim; + let q_dim = config.text.num_attention_heads * head_dim; + let kv_dim = config.text.num_key_value_heads * head_dim; + let intermediate = config.text.intermediate_size; + let mut specs = Vec::with_capacity(2 + 1 + config.text.num_hidden_layers * 11); + specs.push(TensorHeaderSpec { + name: TEXT_EMBEDDING.to_string(), + dtype: "BF16", + shape: vec![config.text.vocab_size, hidden], + }); + specs.push(TensorHeaderSpec { + name: FUSED_MODALITY_EMBEDDING.to_string(), + dtype: "BF16", + shape: fused_modality_shape().to_vec(), + }); + specs.push(TensorHeaderSpec { + name: BODY_NORM.to_string(), + dtype: "BF16", + shape: vec![hidden], + }); + for layer in 0..config.text.num_hidden_layers { + let prefix = format!("body.layers.{layer}"); + for (suffix, shape) in [ + ("input_layernorm.weight", vec![hidden]), + ("post_attention_layernorm.weight", vec![hidden]), + ("self_attn.q_proj.weight", vec![q_dim, hidden]), + ("self_attn.k_proj.weight", vec![kv_dim, hidden]), + ("self_attn.v_proj.weight", vec![kv_dim, hidden]), + ("self_attn.o_proj.weight", vec![hidden, q_dim]), + ("self_attn.q_norm.weight", vec![head_dim]), + ("self_attn.k_norm.weight", vec![head_dim]), + ("mlp.gate_proj.weight", vec![intermediate, hidden]), + ("mlp.up_proj.weight", vec![intermediate, hidden]), + ("mlp.down_proj.weight", vec![hidden, intermediate]), + ] { + specs.push(TensorHeaderSpec { + name: format!("{prefix}.{suffix}"), + dtype: "BF16", + shape, + }); + } + } + specs +} + +pub fn required_body_tensors(config: &HiggsConfig) -> BTreeSet { + let mut names = BTreeSet::new(); + names.insert(BODY_NORM.to_string()); + for layer in 0..config.text.num_hidden_layers { + let prefix = format!("body.layers.{layer}"); + for suffix in [ + "input_layernorm.weight", + "post_attention_layernorm.weight", + "self_attn.q_proj.weight", + "self_attn.k_proj.weight", + "self_attn.v_proj.weight", + "self_attn.o_proj.weight", + "self_attn.q_norm.weight", + "self_attn.k_norm.weight", + "mlp.gate_proj.weight", + "mlp.up_proj.weight", + "mlp.down_proj.weight", + ] { + names.insert(format!("{prefix}.{suffix}")); + } + } + names +} + +pub fn fused_modality_shape() -> [usize; 2] { + [NUM_CODEBOOKS * CODEBOOK_VOCAB_SIZE, HIDDEN_SIZE] +} + +fn is_decoder_only_tensor(name: &str) -> bool { + name.starts_with("tied.embedding.modality_embeddings.0.model.quantizer") + || name.starts_with("tied.embedding.modality_embeddings.0.model.fc2") + || name.starts_with("tied.embedding.modality_embeddings.0.model.acoustic_decoder") +} + +fn read_safetensors_header(path: &Path) -> Result> { + let mut file = std::fs::File::open(path).with_context(|| format!("open {}", path.display()))?; + let mut len_bytes = [0u8; 8]; + file.read_exact(&mut len_bytes) + .with_context(|| format!("read safetensors header length from {}", path.display()))?; + let header_len = usize::try_from(u64::from_le_bytes(len_bytes)).with_context(|| { + format!( + "{} safetensors header length does not fit usize", + path.display() + ) + })?; + ensure!( + header_len < 512 * 1024 * 1024, + "{} safetensors header is unexpectedly large: {} bytes", + path.display(), + header_len + ); + let mut header_bytes = vec![0u8; header_len]; + file.read_exact(&mut header_bytes) + .with_context(|| format!("read safetensors header from {}", path.display()))?; + let value: Value = serde_json::from_slice(&header_bytes) + .with_context(|| format!("parse safetensors header from {}", path.display()))?; + let object = value + .as_object() + .with_context(|| format!("{} safetensors header is not a JSON object", path.display()))?; + let mut tensors = HashMap::new(); + for (name, value) in object { + if name == "__metadata__" { + continue; + } + let dtype = str_field(value, "dtype")?.to_string(); + let shape = value + .get("shape") + .and_then(Value::as_array) + .with_context(|| format!("{name} missing shape"))? + .iter() + .map(|dim| { + dim.as_u64() + .context("shape dim is not u64") + .and_then(|dim| usize::try_from(dim).context("shape dim does not fit usize")) + }) + .collect::>>()?; + let offsets = value + .get("data_offsets") + .and_then(Value::as_array) + .with_context(|| format!("{name} missing data_offsets"))?; + ensure!(offsets.len() == 2, "{name} data_offsets must have length 2"); + let start = offsets[0] + .as_u64() + .with_context(|| format!("{name} data_offsets[0] is not u64")) + .and_then(|offset| usize::try_from(offset).context("offset does not fit usize"))?; + let end = offsets[1] + .as_u64() + .with_context(|| format!("{name} data_offsets[1] is not u64")) + .and_then(|offset| usize::try_from(offset).context("offset does not fit usize"))?; + tensors.insert( + name.clone(), + TensorHeader { + dtype, + shape, + byte_range: [start, end], + }, + ); + } + Ok(tensors) +} + +fn str_field<'a>(value: &'a Value, key: &str) -> Result<&'a str> { + value + .get(key) + .and_then(Value::as_str) + .with_context(|| format!("missing string field {key}")) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::EXPECTED_ARCHITECTURE; + use crate::config::EXPECTED_AUDIO_ENCODER_TYPE; + use crate::config::EXPECTED_HEAD_DIM; + use crate::config::EXPECTED_INTERMEDIATE_SIZE; + use crate::config::EXPECTED_MODEL_TYPE; + use crate::config::EXPECTED_NUM_ATTENTION_HEADS; + use crate::config::EXPECTED_NUM_KV_HEADS; + use crate::config::EXPECTED_NUM_LAYERS; + use crate::config::EXPECTED_ROPE_THETA; + use crate::config::EXPECTED_TEXT_VOCAB_SIZE; + + fn config() -> HiggsConfig { + HiggsConfig::from_json(&serde_json::json!({ + "architectures": [EXPECTED_ARCHITECTURE], + "audio_token_id": -100, + "model_type": EXPECTED_MODEL_TYPE, + "text_config": { + "hidden_size": HIDDEN_SIZE, + "intermediate_size": EXPECTED_INTERMEDIATE_SIZE, + "num_hidden_layers": EXPECTED_NUM_LAYERS, + "num_attention_heads": EXPECTED_NUM_ATTENTION_HEADS, + "num_key_value_heads": EXPECTED_NUM_KV_HEADS, + "head_dim": EXPECTED_HEAD_DIM, + "vocab_size": EXPECTED_TEXT_VOCAB_SIZE, + "rms_norm_eps": 1e-6, + "max_position_embeddings": 32768, + "eos_token_id": 151643, + "tie_word_embeddings": true, + "rope_parameters": {"rope_theta": EXPECTED_ROPE_THETA} + }, + "audio_encoder_config": { + "encoder_type": EXPECTED_AUDIO_ENCODER_TYPE, + "num_codebooks": NUM_CODEBOOKS, + "vocab_size": CODEBOOK_VOCAB_SIZE, + "out_dim": HIDDEN_SIZE, + "tie_word_embeddings": true, + "use_delay_pattern": true + } + })) + .unwrap() + } + + fn manifest_json(include_fused_head: bool) -> Value { + let cfg = config(); + let mut weight_map = serde_json::Map::new(); + weight_map.insert( + TEXT_EMBEDDING.to_string(), + serde_json::json!("model.safetensors"), + ); + if include_fused_head { + weight_map.insert( + FUSED_MODALITY_EMBEDDING.to_string(), + serde_json::json!("model.safetensors"), + ); + } + for name in required_body_tensors(&cfg) { + weight_map.insert(name, serde_json::json!("model.safetensors")); + } + serde_json::json!({"metadata": {"total_size": "8489763794"}, "weight_map": weight_map}) + } + + #[test] + fn validates_required_higgs_manifest_surface() { + let cfg = config(); + let manifest = HiggsWeightManifest::from_json(&manifest_json(true)).unwrap(); + let summary = manifest.validate_for_config(&cfg).unwrap(); + assert_eq!(summary.body_tensors, 397); + assert_eq!(summary.total_tensors, 399); + assert_eq!(fused_modality_shape(), [8208, 2560]); + } + + #[test] + fn rejects_manifest_without_fused_modality_head_weight() { + let cfg = config(); + let manifest = HiggsWeightManifest::from_json(&manifest_json(false)).unwrap(); + let err = manifest.validate_for_config(&cfg).unwrap_err().to_string(); + assert!(err.contains(FUSED_MODALITY_EMBEDDING)); + } + + #[test] + fn expected_checkpoint_tensor_shapes_cover_higgs_body() { + let specs = expected_checkpoint_tensors(&config()); + assert_eq!(specs.len(), 399); + assert!(specs.iter().any(|spec| { + spec.name == "body.layers.0.self_attn.q_proj.weight" + && spec.shape == [4096, HIDDEN_SIZE] + })); + assert!(specs.iter().any(|spec| { + spec.name == "body.layers.0.self_attn.k_proj.weight" + && spec.shape == [1024, HIDDEN_SIZE] + })); + assert!(specs.iter().any(|spec| { + spec.name == "body.layers.0.self_attn.q_norm.weight" + && spec.shape == [EXPECTED_HEAD_DIM] + })); + assert!(specs.iter().any(|spec| { + spec.name == "body.layers.0.mlp.down_proj.weight" + && spec.shape == [HIDDEN_SIZE, EXPECTED_INTERMEDIATE_SIZE] + })); + } + + #[test] + fn validates_checkpoint_header_without_reading_payload() { + let dir = tempfile::TempDir::new().unwrap(); + let cfg = config(); + let manifest = HiggsWeightManifest::from_json(&manifest_json(true)).unwrap(); + write_fake_safetensors_header(dir.path().join("model.safetensors"), &cfg); + + let summary = validate_checkpoint_headers(dir.path(), &cfg, &manifest).unwrap(); + assert_eq!(summary.files_checked, 1); + assert_eq!(summary.tensors_checked, 399); + assert_eq!(summary.bf16_tensors_checked, 399); + } + + #[test] + fn rejects_checkpoint_header_shape_drift() { + let dir = tempfile::TempDir::new().unwrap(); + let cfg = config(); + let manifest = HiggsWeightManifest::from_json(&manifest_json(true)).unwrap(); + let mut specs = expected_checkpoint_tensors(&cfg); + let q_proj = specs + .iter_mut() + .find(|spec| spec.name == "body.layers.0.self_attn.q_proj.weight") + .unwrap(); + q_proj.shape = vec![1, HIDDEN_SIZE]; + write_fake_safetensors_header_from_specs(dir.path().join("model.safetensors"), &specs); + + let err = validate_checkpoint_headers(dir.path(), &cfg, &manifest) + .unwrap_err() + .to_string(); + assert!(err.contains("q_proj")); + assert!(err.contains("shape mismatch")); + } + + fn write_fake_safetensors_header(path: impl AsRef, config: &HiggsConfig) { + write_fake_safetensors_header_from_specs(path, &expected_checkpoint_tensors(config)); + } + + fn write_fake_safetensors_header_from_specs( + path: impl AsRef, + specs: &[TensorHeaderSpec], + ) { + let mut offset = 0usize; + let mut header = serde_json::Map::new(); + for spec in specs { + let bytes = spec.shape.iter().product::() * 2; + header.insert( + spec.name.clone(), + serde_json::json!({ + "dtype": spec.dtype, + "shape": spec.shape, + "data_offsets": [offset, offset + bytes], + }), + ); + offset += bytes; + } + let header = serde_json::to_vec(&Value::Object(header)).unwrap(); + let mut bytes = Vec::new(); + bytes.extend_from_slice(&(header.len() as u64).to_le_bytes()); + bytes.extend_from_slice(&header); + std::fs::write(path, bytes).unwrap(); + } +} diff --git a/pegainfer-qwen3/src/executor.rs b/pegainfer-qwen3/src/executor.rs index 57da546f3..30c0178a6 100644 --- a/pegainfer-qwen3/src/executor.rs +++ b/pegainfer-qwen3/src/executor.rs @@ -13,6 +13,7 @@ use pegainfer_core::kv_pool::KvLayout; use pegainfer_core::ops; use pegainfer_core::tensor::DeviceContext; use pegainfer_core::tensor::HiddenStates; +use pegainfer_core::weight_loader::TensorNameAliases; use pegainfer_core::weight_loader::WeightPrefetch; use pegainfer_core::weight_loader::load_shard_info; use pegainfer_frontend::engine::DeferredFinish; @@ -438,6 +439,36 @@ fn execute_step_on_lane( Ok(WorkerStepOutcome::Ack) } } + StepCommand::PrefillLastHidden { prompt, kv_view } => { + let hidden = lane.execute_prefill_last_hidden(prompt, kv_view)?; + if collect_result { + Ok(WorkerStepOutcome::PrefillHidden(PrefillHiddenResult { + hidden_bf16: hidden, + })) + } else { + Ok(WorkerStepOutcome::Ack) + } + } + StepCommand::PrefillLayerHidden { prompt, kv_view } => { + let hidden = lane.execute_prefill_layer_hidden(prompt, kv_view)?; + if collect_result { + Ok(WorkerStepOutcome::PrefillLayerHidden(hidden)) + } else { + Ok(WorkerStepOutcome::Ack) + } + } + StepCommand::PrefillLayerStages { + layer_idx, + prompt, + kv_view, + } => { + let stages = lane.execute_prefill_layer_stages(*layer_idx, prompt, kv_view)?; + if collect_result { + Ok(WorkerStepOutcome::PrefillStages(stages)) + } else { + Ok(WorkerStepOutcome::Ack) + } + } StepCommand::Unified { prefill_requests, prefill_kv_views, @@ -831,6 +862,29 @@ pub struct UnifiedResult { pub decode_requests: Vec, } +#[derive(Clone, Debug)] +pub struct PrefillHiddenResult { + pub hidden_bf16: Vec, +} + +#[derive(Clone, Debug)] +pub struct RetainedPrefillHiddenResult { + pub request_id: RequestId, + pub hidden_bf16: Vec, +} + +#[derive(Clone, Debug)] +pub struct PrefillLayerHiddenResult { + pub embedding_hidden_bf16: Vec, + pub layer_hidden_bf16: Vec>, + pub final_normed_bf16: Vec, +} + +#[derive(Clone, Debug)] +pub struct PrefillStageResult { + pub stages: Vec<(String, Vec)>, +} + pub(crate) trait ModelExecutor: Send { fn block_size(&self) -> usize; fn max_request_blocks(&self) -> usize; @@ -1230,6 +1284,27 @@ impl Qwen3Executor { ) } + pub fn from_runtime_with_weight_source( + model_path: &str, + weight_path: Option<&str>, + tensor_name_aliases: TensorNameAliases, + enable_cuda_graph: bool, + device_ordinals: &[usize], + ) -> Result { + Self::from_runtime_with_weight_source_and_lora_options( + model_path, + weight_path, + tensor_name_aliases, + enable_cuda_graph, + device_ordinals, + Qwen3LoraOptions::default(), + Qwen3OffloadOptions::disabled(), + crate::scheduler::DEFAULT_MAX_PREFILL_TOKENS, + None, + Qwen3MemoryOptions::default(), + ) + } + #[allow( clippy::needless_pass_by_value, reason = "executor construction is a one-shot ownership boundary" @@ -1243,6 +1318,36 @@ impl Qwen3Executor { max_prefill_tokens: usize, dflash_draft_path: Option<&str>, memory_options: Qwen3MemoryOptions, + ) -> Result { + Self::from_runtime_with_weight_source_and_lora_options( + model_path, + None, + TensorNameAliases::default(), + enable_cuda_graph, + device_ordinals, + lora_options, + offload_options, + max_prefill_tokens, + dflash_draft_path, + memory_options, + ) + } + + #[allow( + clippy::needless_pass_by_value, + reason = "executor construction is a one-shot ownership boundary" + )] + fn from_runtime_with_weight_source_and_lora_options( + model_path: &str, + weight_path: Option<&str>, + tensor_name_aliases: TensorNameAliases, + enable_cuda_graph: bool, + device_ordinals: &[usize], + lora_options: Qwen3LoraOptions, + offload_options: Qwen3OffloadOptions, + max_prefill_tokens: usize, + dflash_draft_path: Option<&str>, + memory_options: Qwen3MemoryOptions, ) -> Result { let mut memory_options = memory_options.validate()?; let lora_options = lora_options.validate()?; @@ -1265,6 +1370,8 @@ impl Qwen3Executor { device_ordinal: device_ordinals[0], max_loras: lora_options.max_loras, max_lora_rank: lora_options.max_lora_rank, + weight_path: weight_path.map(str::to_string), + tensor_name_aliases, }, )?; // The DFlash draft model loads after profiling but lives outside the @@ -1302,7 +1409,8 @@ impl Qwen3Executor { let mut models = Vec::with_capacity(world_size); // TP ranks load sequentially and suppress per-rank prefetch, so keep one // whole-checkpoint prefetch alive across the loop. - let (shard_paths, _) = load_shard_info(model_path)?; + let weight_source = weight_path.unwrap_or(model_path); + let (shard_paths, _) = load_shard_info(weight_source)?; let prefetch = WeightPrefetch::spawn(&shard_paths); for (rank, &device_ordinal) in device_ordinals.iter().enumerate() { models.push(Qwen3Model::from_safetensors_with_runtime( @@ -1313,6 +1421,8 @@ impl Qwen3Executor { device_ordinal, max_loras: lora_options.max_loras, max_lora_rank: lora_options.max_lora_rank, + weight_path: weight_path.map(str::to_string), + tensor_name_aliases: tensor_name_aliases.clone(), }, )?); } @@ -1532,6 +1642,98 @@ impl Qwen3Executor { ::execute_unified(self, plan) } + pub fn prefill_last_hidden_bf16(&mut self, prompt: Vec) -> Result { + self.ensure_single_rank_diagnostic("prefill_last_hidden_bf16")?; + let mut rkv = self.kv_mgr.pool().new_request(prompt.clone(), 1, None); + rkv.schedule_prefill(prompt.len(), self.kv_mgr.pool()) + .map_err(|e| anyhow::anyhow!("diagnostic prefill schedule failed: {e}"))?; + let kv_view = rkv.prefill_view(prompt.len()); + let outcome = self.run_step(&StepCommand::PrefillLastHidden { prompt, kv_view })?; + match outcome { + WorkerStepOutcome::PrefillHidden(result) => Ok(result), + other => Err(anyhow::anyhow!( + "prefill hidden returned unexpected: {}", + other.kind() + )), + } + } + + pub fn prefill_last_hidden_bf16_retained_prompt( + &mut self, + request_id: RequestId, + prompt: Vec, + ) -> Result { + self.ensure_single_rank_diagnostic("prefill_last_hidden_bf16_retained_prompt")?; + anyhow::ensure!( + !self.request_kvs.contains_key(&request_id), + "request {:?} already exists", + request_id + ); + let mut rkv = self.kv_mgr.pool().new_request(prompt.clone(), 1, None); + rkv.schedule_prefill(prompt.len(), self.kv_mgr.pool()) + .map_err(|e| anyhow::anyhow!("diagnostic retained prefill schedule failed: {e}"))?; + let kv_view = rkv.prefill_view(prompt.len()); + let outcome = self.run_step(&StepCommand::PrefillLastHidden { prompt, kv_view })?; + let result = match outcome { + WorkerStepOutcome::PrefillHidden(result) => result, + other => { + return Err(anyhow::anyhow!( + "retained prefill hidden returned unexpected: {}", + other.kind() + )); + } + }; + rkv.apply_prefill_chunk(self.kv_mgr.pool())?; + self.request_kvs.insert(request_id, rkv); + Ok(RetainedPrefillHiddenResult { + request_id, + hidden_bf16: result.hidden_bf16, + }) + } + + pub fn prefill_layer_hidden_bf16( + &mut self, + prompt: Vec, + ) -> Result { + self.ensure_single_rank_diagnostic("prefill_layer_hidden_bf16")?; + let mut rkv = self.kv_mgr.pool().new_request(prompt.clone(), 1, None); + rkv.schedule_prefill(prompt.len(), self.kv_mgr.pool()) + .map_err(|e| anyhow::anyhow!("diagnostic layer prefill schedule failed: {e}"))?; + let kv_view = rkv.prefill_view(prompt.len()); + let outcome = self.run_step(&StepCommand::PrefillLayerHidden { prompt, kv_view })?; + match outcome { + WorkerStepOutcome::PrefillLayerHidden(result) => Ok(result), + other => Err(anyhow::anyhow!( + "prefill layer hidden returned unexpected: {}", + other.kind() + )), + } + } + + pub fn prefill_layer_stages_bf16( + &mut self, + layer_idx: usize, + prompt: Vec, + ) -> Result { + self.ensure_single_rank_diagnostic("prefill_layer_stages_bf16")?; + let mut rkv = self.kv_mgr.pool().new_request(prompt.clone(), 1, None); + rkv.schedule_prefill(prompt.len(), self.kv_mgr.pool()) + .map_err(|e| anyhow::anyhow!("diagnostic stage prefill schedule failed: {e}"))?; + let kv_view = rkv.prefill_view(prompt.len()); + let outcome = self.run_step(&StepCommand::PrefillLayerStages { + layer_idx, + prompt, + kv_view, + })?; + match outcome { + WorkerStepOutcome::PrefillStages(result) => Ok(result), + other => Err(anyhow::anyhow!( + "prefill stages returned unexpected: {}", + other.kind() + )), + } + } + pub fn load_lora_adapter(&mut self, request: &LoadLoraAdapterRequest) -> Result<()> { ::load_lora_adapter(self, request) } @@ -1666,6 +1868,14 @@ impl Qwen3Executor { self.kv_mgr.pool().evict_inactive(); } + fn ensure_single_rank_diagnostic(&self, op: &str) -> Result<()> { + anyhow::ensure!( + self.workers.is_empty(), + "{op} is only supported on the single-GPU diagnostic path" + ); + Ok(()) + } + /// Begin an async CPU-tier KV prefetch for `request_id`; see the /// [`ModelExecutor`] hook. Public so admission drivers and tests can park a /// request on its load. Returns `true` when a load is in flight. @@ -3359,6 +3569,98 @@ impl LocalQwen3Lane { ) } + fn execute_prefill_last_hidden( + &mut self, + prompt: &[u32], + kv_view: &KvView, + ) -> Result> { + let hidden = self.model.prefill_last_normed_hidden( + prompt, + kv_view, + self.kv_buffer.buffer(), + &self.layout, + )?; + let host = self + .model + .device_ctx() + .stream + .clone_dtoh(&hidden.data) + .map_err(|e| anyhow::anyhow!("D2H hidden copy failed: {e}"))?; + self.model.device_ctx().sync()?; + Ok(host) + } + + fn execute_prefill_layer_hidden( + &mut self, + prompt: &[u32], + kv_view: &KvView, + ) -> Result { + let (embedding_hidden, layer_hidden, final_normed) = + self.model.prefill_last_hidden_layer_snapshots( + prompt, + kv_view, + self.kv_buffer.buffer(), + &self.layout, + )?; + let embedding_hidden_bf16 = self + .model + .device_ctx() + .stream + .clone_dtoh(&embedding_hidden.data) + .map_err(|e| anyhow::anyhow!("D2H embedding hidden copy failed: {e}"))?; + let mut layer_hidden_bf16 = Vec::with_capacity(layer_hidden.len()); + for hidden in layer_hidden { + layer_hidden_bf16.push( + self.model + .device_ctx() + .stream + .clone_dtoh(&hidden.data) + .map_err(|e| anyhow::anyhow!("D2H layer hidden copy failed: {e}"))?, + ); + } + let final_normed_bf16 = self + .model + .device_ctx() + .stream + .clone_dtoh(&final_normed.data) + .map_err(|e| anyhow::anyhow!("D2H final normed hidden copy failed: {e}"))?; + self.model.device_ctx().sync()?; + Ok(PrefillLayerHiddenResult { + embedding_hidden_bf16, + layer_hidden_bf16, + final_normed_bf16, + }) + } + + fn execute_prefill_layer_stages( + &mut self, + layer_idx: usize, + prompt: &[u32], + kv_view: &KvView, + ) -> Result { + let stages = self.model.prefill_layer_stage_snapshots( + layer_idx, + prompt, + kv_view, + self.kv_buffer.buffer(), + &self.layout, + )?; + let mut host_stages = Vec::with_capacity(stages.len()); + for stage in stages { + let values = self + .model + .device_ctx() + .stream + .clone_dtoh(&stage.values.data) + .map_err(|e| anyhow::anyhow!("D2H {} copy failed: {e}", stage.name))?; + host_stages.push((stage.name, values)); + } + self.model.device_ctx().sync()?; + Ok(PrefillStageResult { + stages: host_stages, + }) + } + /// DFlash verify forward over each request's `block_size`-token span, using /// the fixed pre-allocated [`VerifyGraphBuffers`] (no per-step allocation). /// Numerically equivalent to the `batch_prefill(echo=true)` verify path it @@ -3506,6 +3808,19 @@ enum StepCommand { kv_views: Vec, sample_seed: u64, }, + PrefillLastHidden { + prompt: Vec, + kv_view: KvView, + }, + PrefillLayerHidden { + prompt: Vec, + kv_view: KvView, + }, + PrefillLayerStages { + layer_idx: usize, + prompt: Vec, + kv_view: KvView, + }, Unified { prefill_requests: Vec, prefill_kv_views: Vec, @@ -3536,7 +3851,9 @@ enum StepCommand { }, /// Speculative draft: roll the DFlash draft model forward one block per /// request. Uses the draft's own KV — no target KV views. - SpeculativeDraft { requests: Vec }, + SpeculativeDraft { + requests: Vec, + }, } impl StepCommand { @@ -3544,6 +3861,9 @@ impl StepCommand { match self { Self::Prefill { .. } => "prefill", Self::Decode { .. } => "decode", + Self::PrefillLastHidden { .. } => "prefill_hidden", + Self::PrefillLayerHidden { .. } => "prefill_layer_hidden", + Self::PrefillLayerStages { .. } => "prefill_layer_stages", Self::Unified { .. } => "unified", Self::SplitConcurrent { .. } => "split_concurrent", Self::SpeculativeVerify { .. } => "speculative_verify", @@ -3627,6 +3947,9 @@ enum WorkerStepOutcome { Prefill(PrefillResult), Decode(DecodeResult), Unified(UnifiedResult), + PrefillHidden(PrefillHiddenResult), + PrefillLayerHidden(PrefillLayerHiddenResult), + PrefillStages(PrefillStageResult), /// Split-concurrent: decode result is ready; prefill is still in-flight /// on the prefill stream. The executor must call a follow-up to sync+sample /// prefill before using prefill scratch buffers again. @@ -3647,6 +3970,9 @@ impl WorkerStepOutcome { Self::Prefill(_) => "prefill", Self::Decode(_) => "decode", Self::Unified(_) => "unified", + Self::PrefillHidden(_) => "prefill_hidden", + Self::PrefillLayerHidden(_) => "prefill_layer_hidden", + Self::PrefillStages(_) => "prefill_layer_stages", Self::SplitDecodeReady { .. } => "split_decode_ready", Self::SpeculativeVerify(_) => "speculative_verify", Self::SpeculativeDraft(_) => "speculative_draft", diff --git a/pegainfer-qwen3/src/lib.rs b/pegainfer-qwen3/src/lib.rs index 59ca74703..4d3ea56f0 100644 --- a/pegainfer-qwen3/src/lib.rs +++ b/pegainfer-qwen3/src/lib.rs @@ -203,12 +203,16 @@ pub mod runtime { pub use crate::executor::DecodeRequestResult; pub use crate::executor::DecodeResult; pub use crate::executor::DecodeStepItem; + pub use crate::executor::PrefillHiddenResult; + pub use crate::executor::PrefillLayerHiddenResult; pub use crate::executor::PrefillPlan; pub use crate::executor::PrefillRequestResult; pub use crate::executor::PrefillResult; + pub use crate::executor::PrefillStageResult; pub use crate::executor::PrefillStepItem; pub use crate::executor::Qwen3Executor; pub use crate::executor::RequestId; + pub use crate::executor::RetainedPrefillHiddenResult; pub use crate::executor::UnifiedPlan; pub use crate::executor::UnifiedResult; } diff --git a/pegainfer-qwen3/src/prefill.rs b/pegainfer-qwen3/src/prefill.rs index 04ce13d92..64a791346 100644 --- a/pegainfer-qwen3/src/prefill.rs +++ b/pegainfer-qwen3/src/prefill.rs @@ -6,6 +6,7 @@ use half::bf16; use pegainfer_core::kv_pool::KvLayout; use pegainfer_core::ops; use pegainfer_core::ops::PrefillPagedPlan; +use pegainfer_core::tensor::DeviceVec; use super::config::PREFILL_ATTENTION_CTA_TILE_Q; use super::weights::Qwen3Model; @@ -99,6 +100,12 @@ pub(super) struct PrefillBuffers { pub(super) attn_output: HiddenStates, // q_dim × seq_len } +/// One named last-token snapshot from a selected prefill layer. +pub(crate) struct PrefillStageSnapshot { + pub(crate) name: String, + pub(crate) values: DeviceVec, +} + impl PrefillBuffers { pub(super) fn new( ctx: &DeviceContext, @@ -141,6 +148,26 @@ impl PrefillBuffers { } impl Qwen3Model { + fn prefill_plan_for_single_prompt( + &self, + prompt: &[u32], + kv_view: &KvView, + ) -> Result { + anyhow::ensure!(!prompt.is_empty(), "prompt must not be empty"); + let start_position = kv_view.seq_len() - prompt.len(); + PrefillPagedPlan::from_raw_batch_with_cta_tile_q( + &self.ctx, + &[kv_view.page_indices().to_vec()], + &[kv_view.last_page_len()], + &[start_position], + &[prompt.len()], + self.local_num_attention_heads(), + self.local_num_key_value_heads(), + self.config.head_dim, + PREFILL_ATTENTION_CTA_TILE_Q, + ) + } + pub(super) fn get_embeddings_batch(&self, token_ids: &[u32]) -> Result { let seq_len = token_ids.len(); let hidden_dim = self.config.hidden_size; @@ -575,6 +602,370 @@ impl Qwen3Model { result } + /// Run one prompt prefill and return the final RMSNorm'ed hidden state for + /// the last prompt token. This narrow diagnostic hook lets model adapters + /// compare the exact hidden-state contract consumed by downstream heads. + pub(crate) fn prefill_last_normed_hidden( + &self, + prompt: &[u32], + kv_view: &KvView, + kv_buffer: &CudaSlice, + layout: &KvLayout, + ) -> Result { + let plan = self.prefill_plan_for_single_prompt(prompt, kv_view)?; + let mut hidden = self.get_embeddings_batch(prompt)?; + let lora_groups: Vec> = Vec::new(); + self.process_all_layers_batch_multi( + &mut hidden, + layout, + kv_buffer, + &plan, + &lora_groups, + None, + )?; + let last_hidden = ops::extract_vec(&self.ctx, &hidden, prompt.len() - 1)?; + let normed = ops::rms_norm( + &self.ctx, + &last_hidden, + &self.norm, + self.config.rms_norm_eps, + )?; + park(hidden); + park(plan); + Ok(normed) + } + + /// Run one prompt prefill and return HuggingFace-style last-token hidden + /// snapshots: embedding output, each transformer block output before final + /// model norm, and the final normed hidden. + pub(crate) fn prefill_last_hidden_layer_snapshots( + &self, + prompt: &[u32], + kv_view: &KvView, + kv_buffer: &CudaSlice, + layout: &KvLayout, + ) -> Result<(DeviceVec, Vec, DeviceVec)> { + let plan = self.prefill_plan_for_single_prompt(prompt, kv_view)?; + let mut hidden = self.get_embeddings_batch(prompt)?; + let embedding_snapshot = ops::extract_vec(&self.ctx, &hidden, prompt.len() - 1)?; + let total_tokens = hidden.seq_len; + let inter_dim = self.local_intermediate_size(); + let q_dim = self.local_q_dim(); + let kv_dim = self.local_kv_dim(); + let last_token_idx = prompt.len() - 1; + let lora_groups: Vec> = Vec::new(); + let mut bufs = PrefillBuffers::new( + &self.ctx, + self.config.hidden_size, + q_dim, + kv_dim, + inter_dim, + total_tokens, + )?; + let mut layer_snapshots = Vec::with_capacity(self.layers.len()); + + crate::green_ctx::fence_producers_before_override(&self.ctx)?; + let run = (|| -> Result<()> { + for (layer_idx, layer) in self.layers.iter().enumerate() { + self.forward_layer_batch_paged( + layer_idx, + layer, + &mut hidden, + kv_buffer, + layout, + &plan, + &lora_groups, + &mut bufs, + )?; + layer_snapshots.push(ops::extract_vec(&self.ctx, &hidden, last_token_idx)?); + } + Ok(()) + })(); + park(bufs); + run?; + + let last_hidden = ops::extract_vec(&self.ctx, &hidden, last_token_idx)?; + let final_normed = ops::rms_norm( + &self.ctx, + &last_hidden, + &self.norm, + self.config.rms_norm_eps, + )?; + park(hidden); + park(plan); + Ok((embedding_snapshot, layer_snapshots, final_normed)) + } + + /// Run layers before `layer_idx` normally, then expose named last-token + /// snapshots across the selected layer's pre-attention, attention, and MLP + /// stages without changing the production layer forward path. + pub(crate) fn prefill_layer_stage_snapshots( + &self, + layer_idx: usize, + prompt: &[u32], + kv_view: &KvView, + kv_buffer: &CudaSlice, + layout: &KvLayout, + ) -> Result> { + anyhow::ensure!( + layer_idx < self.layers.len(), + "layer_idx {} out of range for {} layers", + layer_idx, + self.layers.len() + ); + + let plan = self.prefill_plan_for_single_prompt(prompt, kv_view)?; + let mut hidden = self.get_embeddings_batch(prompt)?; + let total_tokens = hidden.seq_len; + let inter_dim = self.local_intermediate_size(); + let q_dim = self.local_q_dim(); + let kv_dim = self.local_kv_dim(); + let last_token_idx = prompt.len() - 1; + let lora_groups: Vec> = Vec::new(); + let mut stages = Vec::new(); + let mut bufs = PrefillBuffers::new( + &self.ctx, + self.config.hidden_size, + q_dim, + kv_dim, + inter_dim, + total_tokens, + )?; + + crate::green_ctx::fence_producers_before_override(&self.ctx)?; + let run = (|| -> Result<()> { + for (previous_idx, previous_layer) in self.layers.iter().take(layer_idx).enumerate() { + self.forward_layer_batch_paged( + previous_idx, + previous_layer, + &mut hidden, + kv_buffer, + layout, + &plan, + &lora_groups, + &mut bufs, + )?; + } + + let layer = &self.layers[layer_idx]; + let stage_prefix = format!("layer{layer_idx}"); + stages.push(PrefillStageSnapshot { + name: format!("{stage_prefix}.input_hidden.bf16"), + values: ops::extract_vec(&self.ctx, &hidden, last_token_idx)?, + }); + + self.forward_layer_pre_attn(layer_idx, layer, &hidden, &lora_groups, &mut bufs)?; + push_stage( + &mut stages, + format!("{stage_prefix}.input_norm.bf16"), + &self.ctx, + &bufs.normed, + last_token_idx, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.q_proj.bf16"), + &self.ctx, + &bufs.q_batch, + last_token_idx, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.k_proj.bf16"), + &self.ctx, + &bufs.k_batch, + last_token_idx, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.v_proj.bf16"), + &self.ctx, + &bufs.v_batch, + last_token_idx, + )?; + + let mut q_norm_debug = HiddenStates { + data: bufs + .q_batch + .data + .try_clone() + .map_err(|e| anyhow::anyhow!("D2D q norm debug clone failed: {e}"))?, + hidden_dim: bufs.q_batch.hidden_dim, + seq_len: bufs.q_batch.seq_len, + }; + let mut k_norm_debug = HiddenStates { + data: bufs + .k_batch + .data + .try_clone() + .map_err(|e| anyhow::anyhow!("D2D k norm debug clone failed: {e}"))?, + hidden_dim: bufs.k_batch.hidden_dim, + seq_len: bufs.k_batch.seq_len, + }; + let zero_positions = vec![0i32; total_tokens]; + let zero_positions_d = self + .ctx + .stream + .clone_htod(&zero_positions) + .map_err(|e| anyhow::anyhow!("H2D q/k norm debug positions failed: {e}"))?; + ops::qk_norm_rope_batch_decode_into( + &self.ctx, + &mut q_norm_debug, + &mut k_norm_debug, + 0, + total_tokens, + &layer.attention.q_norm, + &layer.attention.k_norm, + &self.cos_cache, + &self.sin_cache, + &zero_positions_d, + self.local_num_attention_heads(), + self.local_num_key_value_heads(), + self.config.head_dim, + self.config.rms_norm_eps, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.q_norm.bf16"), + &self.ctx, + &q_norm_debug, + last_token_idx, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.k_norm.bf16"), + &self.ctx, + &k_norm_debug, + last_token_idx, + )?; + park(zero_positions_d); + park(q_norm_debug); + park(k_norm_debug); + + self.forward_layer_attn(layer_idx, layer, kv_buffer, layout, &plan, &mut bufs)?; + push_stage( + &mut stages, + format!("{stage_prefix}.q_norm_rope.bf16"), + &self.ctx, + &bufs.q_batch, + last_token_idx, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.k_norm_rope.bf16"), + &self.ctx, + &bufs.k_batch, + last_token_idx, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.attn_output.bf16"), + &self.ctx, + &bufs.attn_output, + last_token_idx, + )?; + + ops::gemm_into_checked( + &self.ctx, + &layer.attention.o_proj, + &bufs.attn_output, + &mut bufs.o_buf, + )?; + self.all_reduce_hidden(&mut bufs.o_buf)?; + push_stage( + &mut stages, + format!("{stage_prefix}.o_proj.bf16"), + &self.ctx, + &bufs.o_buf, + last_token_idx, + )?; + + pegainfer_kernels::ops::fused_add_rms_norm_round_batch_into( + &self.ctx, + &mut hidden, + &bufs.o_buf, + &layer.post_attention_layernorm, + self.config.rms_norm_eps, + &mut bufs.normed, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.post_attn_norm.bf16"), + &self.ctx, + &bufs.normed, + last_token_idx, + )?; + + ops::gemm_rows_into_checked( + &self.ctx, + &layer.mlp.gate_up_proj, + 0, + inter_dim, + &bufs.normed, + &mut bufs.gate_out, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.gate_proj.bf16"), + &self.ctx, + &bufs.gate_out, + last_token_idx, + )?; + ops::gemm_rows_into_checked( + &self.ctx, + &layer.mlp.gate_up_proj, + inter_dim, + inter_dim, + &bufs.normed, + &mut bufs.up_out, + )?; + push_stage( + &mut stages, + format!("{stage_prefix}.up_proj.bf16"), + &self.ctx, + &bufs.up_out, + last_token_idx, + )?; + ops::silu_mul_batch_into(&self.ctx, &bufs.gate_out, &bufs.up_out, &mut bufs.act_out)?; + push_stage( + &mut stages, + format!("{stage_prefix}.silu_mul.bf16"), + &self.ctx, + &bufs.act_out, + last_token_idx, + )?; + ops::gemm_into_checked( + &self.ctx, + &layer.mlp.down_proj, + &bufs.act_out, + &mut bufs.o_buf, + )?; + self.all_reduce_hidden(&mut bufs.o_buf)?; + push_stage( + &mut stages, + format!("{stage_prefix}.down_proj.bf16"), + &self.ctx, + &bufs.o_buf, + last_token_idx, + )?; + ops::add_batch_into(&self.ctx, &hidden, &bufs.o_buf, &mut bufs.hidden_out)?; + std::mem::swap(&mut hidden, &mut bufs.hidden_out); + push_stage( + &mut stages, + format!("{stage_prefix}.output_hidden.bf16"), + &self.ctx, + &hidden, + last_token_idx, + )?; + Ok(()) + })(); + park(bufs); + run?; + park(hidden); + park(plan); + Ok(stages) + } + fn process_all_layers_batch_multi( &self, hidden: &mut HiddenStates, @@ -655,3 +1046,17 @@ impl Qwen3Model { Ok(captured_hidden) } } + +fn push_stage( + stages: &mut Vec, + name: String, + ctx: &DeviceContext, + batch: &HiddenStates, + token_idx: usize, +) -> Result<()> { + stages.push(PrefillStageSnapshot { + name, + values: ops::extract_vec(ctx, batch, token_idx)?, + }); + Ok(()) +} diff --git a/pegainfer-qwen3/src/weights/load.rs b/pegainfer-qwen3/src/weights/load.rs index 827a90cba..8c92467cc 100644 --- a/pegainfer-qwen3/src/weights/load.rs +++ b/pegainfer-qwen3/src/weights/load.rs @@ -10,6 +10,7 @@ use pegainfer_core::tensor::DeviceContext; use pegainfer_core::weight_loader::FusedPart; use pegainfer_core::weight_loader::SlotId; use pegainfer_core::weight_loader::StagedWeightLoader; +use pegainfer_core::weight_loader::TensorNameAliases; use pegainfer_core::weight_loader::VecSlotId; use pegainfer_core::weight_loader::WeightPrefetch; use pegainfer_core::weight_loader::deserialize_shards; @@ -25,13 +26,15 @@ use super::TransformerBlock; use crate::config::Config; use crate::config::TensorParallelConfig; -#[derive(Clone, Copy, Debug)] +#[derive(Clone, Debug)] pub(crate) struct ModelRuntimeConfig { pub(crate) enable_cuda_graph: bool, pub(crate) tensor_parallel: Option, pub(crate) device_ordinal: usize, pub(crate) max_loras: usize, pub(crate) max_lora_rank: usize, + pub(crate) weight_path: Option, + pub(crate) tensor_name_aliases: TensorNameAliases, } impl Default for ModelRuntimeConfig { @@ -42,6 +45,8 @@ impl Default for ModelRuntimeConfig { device_ordinal: 0, max_loras: crate::Qwen3LoraOptions::DEFAULT_MAX_LORAS, max_lora_rank: crate::Qwen3LoraOptions::DEFAULT_MAX_LORA_RANK, + weight_path: None, + tensor_name_aliases: TensorNameAliases::default(), } } } @@ -70,7 +75,8 @@ impl Qwen3Model { let tensor_parallel = runtime.tensor_parallel.unwrap_or_default(); tensor_parallel.validate_for(&config)?; - let (shard_paths, weight_map) = load_shard_info(model_path)?; + let weight_path = runtime.weight_path.as_deref().unwrap_or(model_path); + let (shard_paths, weight_map) = load_shard_info(weight_path)?; debug!("Loading {} safetensor shard(s)", shard_paths.len()); let prefetch = (tensor_parallel.world_size == 1).then(|| WeightPrefetch::spawn(&shard_paths)); @@ -78,7 +84,8 @@ impl Qwen3Model { let shards = deserialize_shards(&mmaps)?; let t_gpu = Instant::now(); - let mut loader = StagedWeightLoader::new(&ctx, &shards, &weight_map)?; + let mut loader = StagedWeightLoader::new(&ctx, &shards, &weight_map)? + .with_aliases(runtime.tensor_name_aliases); let hidden = config.hidden_size; debug!("Loading embeddings to GPU"); let embed_slot = loader.matrix("model.embed_tokens.weight", config.vocab_size, hidden)?; diff --git a/test_data/higgs-one-step-audio-logits.safetensors b/test_data/higgs-one-step-audio-logits.safetensors new file mode 100644 index 000000000..90bf54234 Binary files /dev/null and b/test_data/higgs-one-step-audio-logits.safetensors differ diff --git a/tools/accuracy/analyze_higgs_one_step_actual.py b/tools/accuracy/analyze_higgs_one_step_actual.py new file mode 100755 index 000000000..719126ecb --- /dev/null +++ b/tools/accuracy/analyze_higgs_one_step_actual.py @@ -0,0 +1,182 @@ +#!/usr/bin/env python3 +"""Analyze a Higgs Audio one-step PegaInfer actual dump. + +This is a diagnostic companion to `higgs_compare_one_step`: it explains where +the current actual-vs-golden drift comes from instead of only saying pass/fail. +In particular it separates: + +* Qwen3 body hidden-state drift. +* Audio-head implementation drift from CPU fp32 dot vs CUDA bf16 F.linear. +* Top-k instability vs stable argmax decisions. +""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +import torch +from safetensors import safe_open +from safetensors.torch import load_file + +AUDIO_HEAD_KEY = "tied.embedding.modality_embeddings.0.embedding.weight" +CODEBOOKS = 8 +VOCAB = 1026 + + +def _quantile(values: torch.Tensor, q: float) -> float: + if values.numel() == 0: + return 0.0 + return float(torch.quantile(values.float(), q)) + + +def _cosine(left: torch.Tensor, right: torch.Tensor) -> float: + return float( + torch.nn.functional.cosine_similarity( + left.float().flatten(), right.float().flatten(), dim=0 + ) + ) + + +def print_stats(name: str, golden: torch.Tensor, actual: torch.Tensor) -> None: + delta = (actual.float() - golden.float()).flatten() + abs_delta = delta.abs() + print( + f"{name}: " + f"max={float(abs_delta.max()):.6f} " + f"mean={float(abs_delta.mean()):.6f} " + f"p99={_quantile(abs_delta, 0.99):.6f} " + f"rmse={float(torch.sqrt(torch.mean(delta * delta))):.6f} " + f"cos={_cosine(golden, actual):.9f}" + ) + + +def load_audio_head(model_dir: Path) -> torch.Tensor: + shard = model_dir / "model.safetensors" + with safe_open(str(shard), framework="pt", device="cpu") as tensors: + if AUDIO_HEAD_KEY not in tensors.keys(): + raise KeyError(f"missing {AUDIO_HEAD_KEY} in {shard}") + weight = tensors.get_tensor(AUDIO_HEAD_KEY) + expected = (CODEBOOKS * VOCAB, 2560) + if tuple(weight.shape) != expected: + raise ValueError(f"{AUDIO_HEAD_KEY} shape {tuple(weight.shape)} != {expected}") + return weight + + +def topk_overlap(golden_ids: torch.Tensor, actual_ids: torch.Tensor) -> list[int]: + overlaps: list[int] = [] + for cb in range(golden_ids.shape[1]): + golden_set = set(int(v) for v in golden_ids[0, cb].tolist()) + actual_set = set(int(v) for v in actual_ids[0, cb].tolist()) + overlaps.append(len(golden_set & actual_set)) + return overlaps + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--model-dir", type=Path, required=True) + parser.add_argument("--golden", type=Path, required=True) + parser.add_argument("--actual", type=Path, required=True) + parser.add_argument("--device", default="cuda:0") + args = parser.parse_args() + + golden = load_file(str(args.golden), device="cpu") + actual = load_file(str(args.actual), device="cpu") + audio_head = load_audio_head(args.model_dir) + + print("schema:") + for name in ( + "prompt.input_ids_padded", + "prompt.attention_mask", + "prompt.lengths", + "final_hidden.bf16", + "audio_logits.f32", + "audio_argmax.ids", + ): + print( + f" {name}: golden={tuple(golden[name].shape)} {golden[name].dtype} " + f"actual={tuple(actual[name].shape)} {actual[name].dtype}" + ) + + print("\nprimary drift:") + print_stats("final_hidden.bf16", golden["final_hidden.bf16"], actual["final_hidden.bf16"]) + print_stats("audio_logits.f32", golden["audio_logits.f32"], actual["audio_logits.f32"]) + + hidden_delta = ( + actual["final_hidden.bf16"].float().flatten() + - golden["final_hidden.bf16"].float().flatten() + ) + top_idx = torch.topk(hidden_delta.abs(), 12).indices + print("\nfinal_hidden top abs deltas:") + for idx in top_idx.tolist(): + g = float(golden["final_hidden.bf16"].float().flatten()[idx]) + a = float(actual["final_hidden.bf16"].float().flatten()[idx]) + print(f" idx={idx:4d} golden={g:10.6f} actual={a:10.6f} delta={a - g:10.6f}") + + print("\naudio-head dtype attribution:") + golden_hidden = golden["final_hidden.bf16"].to(torch.bfloat16) + actual_hidden = actual["final_hidden.bf16"].to(torch.bfloat16) + golden_logits = golden["audio_logits.f32"].reshape(1, CODEBOOKS, VOCAB) + + cpu_f32_from_golden = (golden_hidden.float() @ audio_head.float().T).reshape( + 1, CODEBOOKS, VOCAB + ) + cpu_f32_from_actual = (actual_hidden.float() @ audio_head.float().T).reshape( + 1, CODEBOOKS, VOCAB + ) + print_stats( + "cpu_f32_from_golden_hidden_vs_golden_logits", + golden_logits, + cpu_f32_from_golden, + ) + print_stats( + "cpu_f32_from_actual_hidden_vs_golden_logits", + golden_logits, + cpu_f32_from_actual, + ) + + if args.device.startswith("cuda") and torch.cuda.is_available(): + weight_cuda = audio_head.to(args.device, dtype=torch.bfloat16) + cuda_bf16_from_golden = torch.nn.functional.linear( + golden_hidden.to(args.device), weight_cuda + ).reshape(1, CODEBOOKS, VOCAB).cpu().float() + cuda_bf16_from_actual = torch.nn.functional.linear( + actual_hidden.to(args.device), weight_cuda + ).reshape(1, CODEBOOKS, VOCAB).cpu().float() + print_stats( + "cuda_bf16_from_golden_hidden_vs_golden_logits", + golden_logits, + cuda_bf16_from_golden, + ) + print_stats( + "cuda_bf16_from_actual_hidden_vs_golden_logits", + golden_logits, + cuda_bf16_from_actual, + ) + print_stats( + "actual_hidden_effect_cuda_bf16", + cuda_bf16_from_golden, + cuda_bf16_from_actual, + ) + else: + print("cuda_bf16 attribution skipped: CUDA is unavailable") + + print("\ntop-k structure:") + overlaps = topk_overlap(golden["audio_top64.ids"], actual["audio_top64.ids"]) + print(f" top64_overlap_by_codebook={overlaps}") + print(f" top64_overlap_min={min(overlaps)} mean={sum(overlaps) / len(overlaps):.2f}") + for cb in range(CODEBOOKS): + golden_row = golden["audio_logits.f32"][0, cb] + actual_row = actual["audio_logits.f32"][0, cb] + top2 = golden_row.topk(2).values + golden_argmax = int(golden_row.argmax()) + actual_argmax = int(actual_row.argmax()) + print( + f" cb={cb} argmax={golden_argmax}/{actual_argmax} " + f"gold_gap={float(top2[0] - top2[1]):.6f} " + f"delta_at_gold_top={float(actual_row[golden_argmax] - golden_row[golden_argmax]):.6f}" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/analyze_higgs_projection_drift.py b/tools/accuracy/analyze_higgs_projection_drift.py new file mode 100644 index 000000000..905b09221 --- /dev/null +++ b/tools/accuracy/analyze_higgs_projection_drift.py @@ -0,0 +1,214 @@ +#!/usr/bin/env python3 +"""Analyze Higgs/Qwen3 projection drift against full layer-stage goldens. + +The trace is an oracle only. This tool recomputes linear projections from +golden and actual stage inputs plus checkpoint weights, then checks whether +PegaInfer's actual GEMM outputs are explained by the same weight/input pair or +whether the projection kernel/storage boundary itself is suspicious. +""" + +from __future__ import annotations + +import argparse +import json +from dataclasses import asdict, dataclass +from pathlib import Path + +import torch +import torch.nn.functional as F +from safetensors import safe_open +from safetensors.torch import load_file + + +@dataclass +class DriftStats: + name: str + shape: list[int] + max_abs: float + mean_abs: float + p99_abs: float + rmse: float + cosine: float + + +@dataclass +class ProjectionReport: + projection: str + input_stage: str + output_stage: str + weight_name: str + input_actual_vs_golden: DriftStats + output_actual_vs_golden: DriftStats + recompute_golden_input_vs_golden_output: DriftStats + recompute_actual_input_vs_actual_output: DriftStats + recompute_actual_input_vs_golden_output: DriftStats + amplification_mean_abs: float + + +PROJECTIONS = { + "q_proj": ("input_norm", "q_proj", "self_attn.q_proj.weight"), + "k_proj": ("input_norm", "k_proj", "self_attn.k_proj.weight"), + "v_proj": ("input_norm", "v_proj", "self_attn.v_proj.weight"), + "o_proj": ("attn_output", "o_proj", "self_attn.o_proj.weight"), + "gate_proj": ("post_attn_norm", "gate_proj", "mlp.gate_proj.weight"), + "up_proj": ("post_attn_norm", "up_proj", "mlp.up_proj.weight"), + "down_proj": ("silu_mul", "down_proj", "mlp.down_proj.weight"), +} + + +def stats(name: str, left: torch.Tensor, right: torch.Tensor) -> DriftStats: + left = left.float().cpu() + right = right.float().cpu() + delta = left.flatten() - right.flatten() + abs_delta = delta.abs() + left_flat = left.flatten() + right_flat = right.flatten() + if left_flat.numel() == 0: + cosine = 1.0 + elif float(torch.linalg.vector_norm(left_flat)) == 0.0 and float( + torch.linalg.vector_norm(right_flat) + ) == 0.0: + cosine = 1.0 + else: + cosine = float(torch.nn.functional.cosine_similarity(left_flat, right_flat, dim=0)) + return DriftStats( + name=name, + shape=list(left.shape), + max_abs=float(abs_delta.max()) if abs_delta.numel() else 0.0, + mean_abs=float(abs_delta.mean()) if abs_delta.numel() else 0.0, + p99_abs=float(torch.quantile(abs_delta, 0.99)) if abs_delta.numel() else 0.0, + rmse=float(torch.sqrt(torch.mean(delta * delta))) if delta.numel() else 0.0, + cosine=cosine, + ) + + +def checkpoint_weight(model_file: Path, layer_idx: int, suffix: str, device: str) -> tuple[str, torch.Tensor]: + key = f"body.layers.{layer_idx}.{suffix}" + with safe_open(str(model_file), framework="pt", device="cpu") as reader: + if key not in reader.keys(): + raise KeyError(f"checkpoint is missing {key}") + weight = reader.get_tensor(key).to(torch.bfloat16).contiguous() + return key, weight.to(device=device) + + +def stage_name(layer_idx: int, suffix: str) -> str: + return f"layer{layer_idx}.{suffix}.bf16" + + +def linear_bf16(input_row: torch.Tensor, weight: torch.Tensor, device: str) -> torch.Tensor: + x = input_row.to(device=device, dtype=torch.bfloat16).contiguous() + y = F.linear(x, weight) + return y.to(torch.bfloat16).cpu().contiguous() + + +def analyze_projection( + name: str, + *, + layer_idx: int, + golden: dict[str, torch.Tensor], + actual: dict[str, torch.Tensor], + model_file: Path, + device: str, +) -> ProjectionReport: + input_suffix, output_suffix, weight_suffix = PROJECTIONS[name] + input_stage = stage_name(layer_idx, input_suffix) + output_stage = stage_name(layer_idx, output_suffix) + weight_name, weight = checkpoint_weight(model_file, layer_idx, weight_suffix, device) + + golden_input = golden[input_stage].cpu().to(torch.bfloat16).contiguous() + actual_input = actual[input_stage].cpu().to(torch.bfloat16).contiguous() + golden_output = golden[output_stage].cpu().to(torch.bfloat16).contiguous() + actual_output = actual[output_stage].cpu().to(torch.bfloat16).contiguous() + + recompute_golden = linear_bf16(golden_input, weight, device) + recompute_actual = linear_bf16(actual_input, weight, device) + + input_drift = stats(f"{name}.input_actual_vs_golden", actual_input, golden_input) + output_drift = stats(f"{name}.output_actual_vs_golden", actual_output, golden_output) + input_mean = max(input_drift.mean_abs, 1e-12) + return ProjectionReport( + projection=name, + input_stage=input_stage, + output_stage=output_stage, + weight_name=weight_name, + input_actual_vs_golden=input_drift, + output_actual_vs_golden=output_drift, + recompute_golden_input_vs_golden_output=stats( + f"{name}.recompute_golden_input_vs_golden_output", + recompute_golden, + golden_output, + ), + recompute_actual_input_vs_actual_output=stats( + f"{name}.recompute_actual_input_vs_actual_output", + recompute_actual, + actual_output, + ), + recompute_actual_input_vs_golden_output=stats( + f"{name}.recompute_actual_input_vs_golden_output", + recompute_actual, + golden_output, + ), + amplification_mean_abs=output_drift.mean_abs / input_mean, + ) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--golden", required=True, help="full layer-stage golden safetensors") + parser.add_argument("--actual", required=True, help="PegaInfer layer-stage actual safetensors") + parser.add_argument("--model-safetensors", required=True) + parser.add_argument("--layer-idx", type=int, required=True) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--projection", action="append", choices=sorted(PROJECTIONS)) + parser.add_argument("--json-out", default="") + args = parser.parse_args() + + if args.device.startswith("cuda") and not torch.cuda.is_available(): + raise RuntimeError(f"{args.device} requested but torch.cuda.is_available() is false") + + golden = load_file(args.golden) + actual = load_file(args.actual) + selected = args.projection or list(PROJECTIONS) + reports = [ + analyze_projection( + projection, + layer_idx=args.layer_idx, + golden=golden, + actual=actual, + model_file=Path(args.model_safetensors), + device=args.device, + ) + for projection in selected + ] + + payload = { + "golden": args.golden, + "actual": args.actual, + "model_safetensors": args.model_safetensors, + "layer_idx": args.layer_idx, + "device": args.device, + "reports": [asdict(report) for report in reports], + } + if args.json_out: + out = Path(args.json_out) + out.parent.mkdir(parents=True, exist_ok=True) + out.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8") + + print(f"layer={args.layer_idx} device={args.device}") + print( + "projection input_mean output_mean recompute_actual_vs_actual " + "recompute_golden_vs_golden amp" + ) + for report in reports: + print( + f"{report.projection:18} " + f"{report.input_actual_vs_golden.mean_abs:10.8f} " + f"{report.output_actual_vs_golden.mean_abs:11.8f} " + f"{report.recompute_actual_input_vs_actual_output.mean_abs:26.8f} " + f"{report.recompute_golden_input_vs_golden_output.mean_abs:26.8f} " + f"{report.amplification_mean_abs:6.2f}" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/analyze_higgs_qk_norm_drift.py b/tools/accuracy/analyze_higgs_qk_norm_drift.py new file mode 100644 index 000000000..beb0cb82d --- /dev/null +++ b/tools/accuracy/analyze_higgs_qk_norm_drift.py @@ -0,0 +1,249 @@ +#!/usr/bin/env python3 +"""Analyze Higgs/Qwen3 q/k RMSNorm drift against a rich trace golden. + +This is a diagnostic tool, not a fixer. It uses the trace as an oracle and +recomputes q/k RMSNorm from the dumped q/k projection tensors plus checkpoint +weights to separate: + +1. projection/input drift that is amplified by RMSNorm, from +2. an actual implementation mismatch in the q/k RMSNorm kernel. +""" + +from __future__ import annotations + +import argparse +import json +from dataclasses import asdict, dataclass +from pathlib import Path + +import torch +from safetensors import safe_open +from safetensors.torch import load_file + + +@dataclass +class DriftStats: + name: str + shape: list[int] + max_abs: float + mean_abs: float + p99_abs: float + rmse: float + cosine: float + + +@dataclass +class VariantReport: + tensor: str + variant: str + recompute_from_golden_proj_vs_golden_norm: DriftStats + recompute_from_actual_proj_vs_actual_norm: DriftStats + recompute_from_actual_proj_vs_golden_norm: DriftStats + + +def stats(name: str, left: torch.Tensor, right: torch.Tensor) -> DriftStats: + left = left.float() + right = right.float() + delta = left.flatten() - right.flatten() + abs_delta = delta.abs() + left_flat = left.flatten() + right_flat = right.flatten() + if left_flat.numel() == 0: + cosine = 1.0 + elif float(torch.linalg.vector_norm(left_flat)) == 0.0 and float( + torch.linalg.vector_norm(right_flat) + ) == 0.0: + cosine = 1.0 + else: + cosine = float(torch.nn.functional.cosine_similarity(left_flat, right_flat, dim=0)) + return DriftStats( + name=name, + shape=list(left.shape), + max_abs=float(abs_delta.max()) if abs_delta.numel() else 0.0, + mean_abs=float(abs_delta.mean()) if abs_delta.numel() else 0.0, + p99_abs=float(torch.quantile(abs_delta, 0.99)) if abs_delta.numel() else 0.0, + rmse=float(torch.sqrt(torch.mean(delta * delta))) if delta.numel() else 0.0, + cosine=cosine, + ) + + +def last_token_from_trace( + tensors: dict[str, torch.Tensor], + name: str, + actual_shape: torch.Size, +) -> torch.Tensor: + tensor = tensors[name] + if tuple(tensor.shape) == tuple(actual_shape): + return tensor + lengths = tensors["prompt.lengths"] + if tensor.ndim >= 3: + rows = [tensor[row, int(length) - 1] for row, length in enumerate(lengths.tolist())] + last = torch.stack(rows, dim=0) + if tuple(last.shape) == tuple(actual_shape): + return last + flat = last.reshape(last.shape[0], -1) + if tuple(flat.shape) == tuple(actual_shape): + return flat + raise ValueError(f"cannot align {name}: golden={tuple(tensor.shape)} actual={tuple(actual_shape)}") + + +def checkpoint_weight(model_file: Path, layer_idx: int, q_or_k: str) -> torch.Tensor: + key = f"body.layers.{layer_idx}.self_attn.{q_or_k}_norm.weight" + with safe_open(str(model_file), framework="pt", device="cpu") as reader: + if key not in reader.keys(): + raise KeyError(f"checkpoint is missing {key}") + return reader.get_tensor(key).to(torch.bfloat16).contiguous() + + +def rmsnorm_variants(x: torch.Tensor, weight: torch.Tensor, head_dim: int, eps: float) -> dict[str, torch.Tensor]: + x = x.reshape(x.shape[0], -1, head_dim) + w_bf16 = weight.to(torch.bfloat16).reshape(1, 1, head_dim) + w_f32 = weight.float().reshape(1, 1, head_dim) + xf = x.float() + inv_rms = torch.rsqrt(torch.mean(xf * xf, dim=-1, keepdim=True) + eps) + norm_f32 = xf * inv_rms + return { + # Mirrors Transformers Qwen3RMSNorm: normalize in fp32, cast to input dtype, + # then multiply by bf16 weight. PyTorch bf16 * bf16 yields bf16. + "hf_like_bf16_mid": (norm_f32.to(torch.bfloat16) * w_bf16) + .to(torch.bfloat16) + .reshape(x.shape[0], -1), + # One final round only. If this wins, the CUDA kernel is over-rounding. + "single_round_fp32_weight": (norm_f32 * w_f32).to(torch.bfloat16).reshape(x.shape[0], -1), + # Round normalized activations first, but multiply with fp32 weight. + "bf16_mid_fp32_weight": (norm_f32.to(torch.bfloat16).float() * w_f32) + .to(torch.bfloat16) + .reshape(x.shape[0], -1), + # Multiply before casting either operand back to bf16. + "fp32_no_mid_round": (norm_f32 * w_bf16.float()).to(torch.bfloat16).reshape(x.shape[0], -1), + } + + +def head_mean_abs(left: torch.Tensor, right: torch.Tensor, head_dim: int) -> list[float]: + delta = (left.float() - right.float()).abs().reshape(left.shape[0], -1, head_dim) + return [float(x) for x in delta.mean(dim=(0, 2)).tolist()] + + +def analyze_tensor( + tensor: str, + *, + golden: dict[str, torch.Tensor], + actual: dict[str, torch.Tensor], + weight: torch.Tensor, + layer_idx: int, + head_dim: int, + eps: float, +) -> tuple[list[VariantReport], dict[str, object]]: + proj_actual_name = f"layer{layer_idx}.{tensor}_proj.bf16" + norm_actual_name = f"layer{layer_idx}.{tensor}_norm.bf16" + proj_trace_name = f"layer.{layer_idx:02}.self_attn.{tensor}_proj.output.bf16" + norm_trace_name = f"layer.{layer_idx:02}.self_attn.{tensor}_norm.output.bf16" + + actual_proj = actual[proj_actual_name].cpu().to(torch.bfloat16).contiguous() + actual_norm = actual[norm_actual_name].cpu().to(torch.bfloat16).contiguous() + golden_proj = last_token_from_trace(golden, proj_trace_name, actual_proj.shape).cpu().to(torch.bfloat16) + golden_norm = last_token_from_trace(golden, norm_trace_name, actual_norm.shape).cpu().to(torch.bfloat16) + + reports = [] + golden_variants = rmsnorm_variants(golden_proj, weight, head_dim, eps) + actual_variants = rmsnorm_variants(actual_proj, weight, head_dim, eps) + for variant in sorted(golden_variants): + reports.append( + VariantReport( + tensor=tensor, + variant=variant, + recompute_from_golden_proj_vs_golden_norm=stats( + f"{tensor}.{variant}.golden_proj_vs_golden_norm", + golden_variants[variant], + golden_norm, + ), + recompute_from_actual_proj_vs_actual_norm=stats( + f"{tensor}.{variant}.actual_proj_vs_actual_norm", + actual_variants[variant], + actual_norm, + ), + recompute_from_actual_proj_vs_golden_norm=stats( + f"{tensor}.{variant}.actual_proj_vs_golden_norm", + actual_variants[variant], + golden_norm, + ), + ) + ) + + direct = { + "proj_actual_vs_golden": asdict(stats(f"{tensor}.proj_actual_vs_golden", actual_proj, golden_proj)), + "norm_actual_vs_golden": asdict(stats(f"{tensor}.norm_actual_vs_golden", actual_norm, golden_norm)), + "norm_head_mean_abs": head_mean_abs(actual_norm, golden_norm, head_dim), + "proj_head_mean_abs": head_mean_abs(actual_proj, golden_proj, head_dim), + } + return reports, direct + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--golden", required=True, help="rich trace golden safetensors") + parser.add_argument("--actual", required=True, help="PegaInfer layer-stage dump safetensors") + parser.add_argument("--model-safetensors", required=True, help="Higgs model.safetensors") + parser.add_argument("--layer-idx", type=int, required=True) + parser.add_argument("--head-dim", type=int, default=128) + parser.add_argument("--eps", type=float, default=1e-6) + parser.add_argument("--json-out", default="") + args = parser.parse_args() + + golden = load_file(args.golden) + actual = load_file(args.actual) + all_reports: list[VariantReport] = [] + direct: dict[str, object] = {} + for tensor in ("q", "k"): + reports, tensor_direct = analyze_tensor( + tensor, + golden=golden, + actual=actual, + weight=checkpoint_weight(Path(args.model_safetensors), args.layer_idx, tensor), + layer_idx=args.layer_idx, + head_dim=args.head_dim, + eps=args.eps, + ) + all_reports.extend(reports) + direct[tensor] = tensor_direct + + payload = { + "golden": args.golden, + "actual": args.actual, + "model_safetensors": args.model_safetensors, + "layer_idx": args.layer_idx, + "head_dim": args.head_dim, + "eps": args.eps, + "direct": direct, + "variants": [asdict(report) for report in all_reports], + } + if args.json_out: + out = Path(args.json_out) + out.parent.mkdir(parents=True, exist_ok=True) + out.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8") + + print(f"layer={args.layer_idx} head_dim={args.head_dim} eps={args.eps}") + for tensor, tensor_direct in direct.items(): + proj = tensor_direct["proj_actual_vs_golden"] + norm = tensor_direct["norm_actual_vs_golden"] + print( + f"{tensor}: proj_mean={proj['mean_abs']:.8f} norm_mean={norm['mean_abs']:.8f} " + f"proj_p99={proj['p99_abs']:.8f} norm_p99={norm['p99_abs']:.8f}" + ) + print("variant ranking by actual_proj_vs_actual_norm mean_abs:") + for report in sorted( + all_reports, + key=lambda item: item.recompute_from_actual_proj_vs_actual_norm.mean_abs, + ): + a = report.recompute_from_actual_proj_vs_actual_norm + g = report.recompute_from_golden_proj_vs_golden_norm + propagated = report.recompute_from_actual_proj_vs_golden_norm + print( + f"{report.tensor}.{report.variant}: " + f"actual_kernel_mean={a.mean_abs:.8f} golden_formula_mean={g.mean_abs:.8f} " + f"propagated_vs_golden_mean={propagated.mean_abs:.8f}" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/analyze_higgs_residual_drift.py b/tools/accuracy/analyze_higgs_residual_drift.py new file mode 100644 index 000000000..dd904cab4 --- /dev/null +++ b/tools/accuracy/analyze_higgs_residual_drift.py @@ -0,0 +1,274 @@ +#!/usr/bin/env python3 +"""Analyze Higgs/Qwen3 residual add and fused-add-RMSNorm drift. + +This diagnostic checks the two residual boundaries in a decoder layer: + +1. input_hidden + o_proj -> post_attn_norm +2. rounded_attn_residual + down_proj -> output_hidden + +The trace/golden tensors remain oracles only; recomputation uses the recorded +stage inputs and checkpoint weights to explain drift without injecting golden +values into runtime code. +""" + +from __future__ import annotations + +import argparse +import json +from dataclasses import asdict, dataclass +from pathlib import Path + +import torch +from safetensors import safe_open +from safetensors.torch import load_file + + +@dataclass +class DriftStats: + name: str + shape: list[int] + max_abs: float + mean_abs: float + p99_abs: float + rmse: float + cosine: float + + +@dataclass +class VariantReport: + variant: str + recompute_golden_post_attn_norm_vs_golden: DriftStats + recompute_actual_post_attn_norm_vs_actual: DriftStats + recompute_actual_post_attn_norm_vs_golden: DriftStats + + +def stats(name: str, left: torch.Tensor, right: torch.Tensor) -> DriftStats: + left = left.float().cpu() + right = right.float().cpu() + delta = left.flatten() - right.flatten() + abs_delta = delta.abs() + left_flat = left.flatten() + right_flat = right.flatten() + if left_flat.numel() == 0: + cosine = 1.0 + elif float(torch.linalg.vector_norm(left_flat)) == 0.0 and float( + torch.linalg.vector_norm(right_flat) + ) == 0.0: + cosine = 1.0 + else: + cosine = float(torch.nn.functional.cosine_similarity(left_flat, right_flat, dim=0)) + return DriftStats( + name=name, + shape=list(left.shape), + max_abs=float(abs_delta.max()) if abs_delta.numel() else 0.0, + mean_abs=float(abs_delta.mean()) if abs_delta.numel() else 0.0, + p99_abs=float(torch.quantile(abs_delta, 0.99)) if abs_delta.numel() else 0.0, + rmse=float(torch.sqrt(torch.mean(delta * delta))) if delta.numel() else 0.0, + cosine=cosine, + ) + + +def stage(layer_idx: int, suffix: str) -> str: + return f"layer{layer_idx}.{suffix}.bf16" + + +def checkpoint_vec(model_file: Path, layer_idx: int, suffix: str, device: str) -> torch.Tensor: + key = f"body.layers.{layer_idx}.{suffix}" + with safe_open(str(model_file), framework="pt", device="cpu") as reader: + if key not in reader.keys(): + raise KeyError(f"checkpoint is missing {key}") + value = reader.get_tensor(key).to(torch.bfloat16).contiguous() + return value.to(device=device) + + +def rmsnorm_variants(x_bf16: torch.Tensor, weight_bf16: torch.Tensor, eps: float) -> dict[str, torch.Tensor]: + x = x_bf16.float() + w_bf16 = weight_bf16.to(torch.bfloat16) + w_f32 = weight_bf16.float() + inv_rms = torch.rsqrt(torch.mean(x * x, dim=-1, keepdim=True) + eps) + norm_f32 = x * inv_rms + return { + # Mirrors Transformers Qwen3RMSNorm. + "hf_like_bf16_mid": (norm_f32.to(torch.bfloat16) * w_bf16).to(torch.bfloat16), + # Mirrors pegainfer::norm::FusedAddRMSNormRoundKernel's visible formula: + # rounded add feeds RMS reduction, then norm * weight is rounded at store. + "fused_round_formula": (norm_f32 * w_f32).to(torch.bfloat16), + "bf16_mid_fp32_weight": (norm_f32.to(torch.bfloat16).float() * w_f32).to(torch.bfloat16), + } + + +def residual_sum_variants(left: torch.Tensor, right: torch.Tensor) -> dict[str, torch.Tensor]: + left = left.to(torch.bfloat16) + right = right.to(torch.bfloat16) + return { + "bf16_add": (left + right).to(torch.bfloat16), + "fp32_add_then_bf16": (left.float() + right.float()).to(torch.bfloat16), + } + + +def analyze_side( + prefix: str, + tensors: dict[str, torch.Tensor], + *, + layer_idx: int, + post_weight: torch.Tensor, + eps: float, + device: str, +) -> dict[str, torch.Tensor | dict[str, torch.Tensor]]: + input_hidden = tensors[stage(layer_idx, "input_hidden")].to(device=device, dtype=torch.bfloat16) + o_proj = tensors[stage(layer_idx, "o_proj")].to(device=device, dtype=torch.bfloat16) + down_proj = tensors[stage(layer_idx, "down_proj")].to(device=device, dtype=torch.bfloat16) + + residual_variants = residual_sum_variants(input_hidden, o_proj) + post_norm_variants = {} + output_variants = {} + for add_name, attn_residual in residual_variants.items(): + for norm_name, post_norm in rmsnorm_variants(attn_residual, post_weight, eps).items(): + post_norm_variants[f"{add_name}+{norm_name}"] = post_norm.cpu().contiguous() + for mlp_add_name, output_hidden in residual_sum_variants(attn_residual, down_proj).items(): + output_variants[f"{add_name}+{mlp_add_name}"] = output_hidden.cpu().contiguous() + return { + f"{prefix}_attn_residual": {k: v.cpu().contiguous() for k, v in residual_variants.items()}, + f"{prefix}_post_norm": post_norm_variants, + f"{prefix}_output_hidden": output_variants, + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--golden", required=True, help="full layer-stage golden safetensors") + parser.add_argument("--actual", required=True, help="PegaInfer layer-stage actual safetensors") + parser.add_argument("--model-safetensors", required=True) + parser.add_argument("--layer-idx", type=int, required=True) + parser.add_argument("--eps", type=float, default=1e-6) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--json-out", default="") + args = parser.parse_args() + + if args.device.startswith("cuda") and not torch.cuda.is_available(): + raise RuntimeError(f"{args.device} requested but torch.cuda.is_available() is false") + + golden = load_file(args.golden) + actual = load_file(args.actual) + post_weight = checkpoint_vec( + Path(args.model_safetensors), + args.layer_idx, + "post_attention_layernorm.weight", + args.device, + ) + golden_recompute = analyze_side( + "golden", golden, layer_idx=args.layer_idx, post_weight=post_weight, eps=args.eps, device=args.device + ) + actual_recompute = analyze_side( + "actual", actual, layer_idx=args.layer_idx, post_weight=post_weight, eps=args.eps, device=args.device + ) + golden_post = golden[stage(args.layer_idx, "post_attn_norm")].cpu().to(torch.bfloat16) + actual_post = actual[stage(args.layer_idx, "post_attn_norm")].cpu().to(torch.bfloat16) + golden_output = golden[stage(args.layer_idx, "output_hidden")].cpu().to(torch.bfloat16) + actual_output = actual[stage(args.layer_idx, "output_hidden")].cpu().to(torch.bfloat16) + + post_reports = [] + for variant in sorted(golden_recompute["golden_post_norm"]): + post_reports.append( + VariantReport( + variant=variant, + recompute_golden_post_attn_norm_vs_golden=stats( + f"{variant}.golden_post_norm_vs_golden", + golden_recompute["golden_post_norm"][variant], + golden_post, + ), + recompute_actual_post_attn_norm_vs_actual=stats( + f"{variant}.actual_post_norm_vs_actual", + actual_recompute["actual_post_norm"][variant], + actual_post, + ), + recompute_actual_post_attn_norm_vs_golden=stats( + f"{variant}.actual_post_norm_vs_golden", + actual_recompute["actual_post_norm"][variant], + golden_post, + ), + ) + ) + + output_reports = [] + for variant in sorted(golden_recompute["golden_output_hidden"]): + output_reports.append( + { + "variant": variant, + "recompute_golden_output_vs_golden": asdict( + stats( + f"{variant}.golden_output_vs_golden", + golden_recompute["golden_output_hidden"][variant], + golden_output, + ) + ), + "recompute_actual_output_vs_actual": asdict( + stats( + f"{variant}.actual_output_vs_actual", + actual_recompute["actual_output_hidden"][variant], + actual_output, + ) + ), + "recompute_actual_output_vs_golden": asdict( + stats( + f"{variant}.actual_output_vs_golden", + actual_recompute["actual_output_hidden"][variant], + golden_output, + ) + ), + } + ) + + direct = { + "o_proj_actual_vs_golden": asdict( + stats(stage(args.layer_idx, "o_proj"), actual[stage(args.layer_idx, "o_proj")], golden[stage(args.layer_idx, "o_proj")]) + ), + "post_attn_norm_actual_vs_golden": asdict(stats(stage(args.layer_idx, "post_attn_norm"), actual_post, golden_post)), + "down_proj_actual_vs_golden": asdict( + stats(stage(args.layer_idx, "down_proj"), actual[stage(args.layer_idx, "down_proj")], golden[stage(args.layer_idx, "down_proj")]) + ), + "output_hidden_actual_vs_golden": asdict(stats(stage(args.layer_idx, "output_hidden"), actual_output, golden_output)), + } + payload = { + "golden": args.golden, + "actual": args.actual, + "model_safetensors": args.model_safetensors, + "layer_idx": args.layer_idx, + "eps": args.eps, + "device": args.device, + "direct": direct, + "post_attn_norm_variants": [asdict(report) for report in post_reports], + "output_hidden_variants": output_reports, + } + if args.json_out: + out = Path(args.json_out) + out.parent.mkdir(parents=True, exist_ok=True) + out.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8") + + print(f"layer={args.layer_idx} device={args.device} eps={args.eps}") + print( + "post_variant golden_vs_golden actual_vs_actual " + "actual_vs_golden" + ) + for report in sorted( + post_reports, + key=lambda item: item.recompute_actual_post_attn_norm_vs_actual.mean_abs, + ): + print( + f"{report.variant:36} " + f"{report.recompute_golden_post_attn_norm_vs_golden.mean_abs:16.8f} " + f"{report.recompute_actual_post_attn_norm_vs_actual.mean_abs:16.8f} " + f"{report.recompute_actual_post_attn_norm_vs_golden.mean_abs:16.8f}" + ) + print("output_variant golden_vs_golden actual_vs_actual actual_vs_golden") + for report in sorted(output_reports, key=lambda item: item["recompute_actual_output_vs_actual"]["mean_abs"]): + print( + f"{report['variant']:36} " + f"{report['recompute_golden_output_vs_golden']['mean_abs']:16.8f} " + f"{report['recompute_actual_output_vs_actual']['mean_abs']:16.8f} " + f"{report['recompute_actual_output_vs_golden']['mean_abs']:16.8f}" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/analyze_higgs_rmsnorm_drift.py b/tools/accuracy/analyze_higgs_rmsnorm_drift.py new file mode 100644 index 000000000..c6aabf3d9 --- /dev/null +++ b/tools/accuracy/analyze_higgs_rmsnorm_drift.py @@ -0,0 +1,213 @@ +#!/usr/bin/env python3 +"""Analyze Higgs/Qwen3 plain RMSNorm drift. + +This diagnostic checks RMSNorm boundaries that are not covered by the +fused-add-RMSNorm helper: + +* decoder layer input RMSNorm from a full-stage layer dump +* final model RMSNorm from the last decoder output to final_hidden + +The golden tensors stay read-only oracles. Recompute variants use dumped inputs +and checkpoint weights to decide whether drift is explained by input drift or by +a real RMSNorm rounding semantic mismatch. +""" + +from __future__ import annotations + +import argparse +import json +from dataclasses import asdict, dataclass +from pathlib import Path + +import torch +from safetensors import safe_open +from safetensors.torch import load_file + + +@dataclass +class DriftStats: + name: str + shape: list[int] + max_abs: float + mean_abs: float + p99_abs: float + rmse: float + cosine: float + + +def stats(name: str, left: torch.Tensor, right: torch.Tensor) -> DriftStats: + left = left.float().cpu() + right = right.float().cpu() + delta = left.flatten() - right.flatten() + abs_delta = delta.abs() + if left.numel() == 0: + cosine = 1.0 + elif float(torch.linalg.vector_norm(left.flatten())) == 0.0 and float( + torch.linalg.vector_norm(right.flatten()) + ) == 0.0: + cosine = 1.0 + else: + cosine = float(torch.nn.functional.cosine_similarity(left.flatten(), right.flatten(), dim=0)) + return DriftStats( + name=name, + shape=list(left.shape), + max_abs=float(abs_delta.max()) if abs_delta.numel() else 0.0, + mean_abs=float(abs_delta.mean()) if abs_delta.numel() else 0.0, + p99_abs=float(torch.quantile(abs_delta, 0.99)) if abs_delta.numel() else 0.0, + rmse=float(torch.sqrt(torch.mean(delta * delta))) if delta.numel() else 0.0, + cosine=cosine, + ) + + +def checkpoint_vec(model_file: Path, key: str, device: str) -> torch.Tensor: + with safe_open(str(model_file), framework="pt", device="cpu") as reader: + if key not in reader.keys(): + raise KeyError(f"checkpoint is missing {key}") + value = reader.get_tensor(key).to(torch.bfloat16).contiguous() + return value.to(device=device) + + +def rmsnorm_variants(x_bf16: torch.Tensor, weight_bf16: torch.Tensor, eps: float) -> dict[str, torch.Tensor]: + x = x_bf16.to(torch.bfloat16) + w_bf16 = weight_bf16.to(torch.bfloat16) + w_f32 = weight_bf16.float() + xf = x.float() + inv_rms = torch.rsqrt(torch.mean(xf * xf, dim=-1, keepdim=True) + eps) + norm_f32 = xf * inv_rms + return { + "hf_like_bf16_mid": (norm_f32.to(torch.bfloat16) * w_bf16).to(torch.bfloat16), + "single_round_fp32_weight": (norm_f32 * w_f32).to(torch.bfloat16), + "bf16_mid_fp32_weight": (norm_f32.to(torch.bfloat16).float() * w_f32).to(torch.bfloat16), + } + + +def add_variant_reports( + payload: dict[str, object], + *, + boundary: str, + golden_input: torch.Tensor, + golden_output: torch.Tensor, + actual_input: torch.Tensor, + actual_output: torch.Tensor, + weight: torch.Tensor, + eps: float, +) -> None: + reports = [] + golden_variants = rmsnorm_variants(golden_input, weight, eps) + actual_variants = rmsnorm_variants(actual_input, weight, eps) + for variant in sorted(golden_variants): + reports.append( + { + "variant": variant, + "golden_recompute_vs_golden_output": asdict( + stats(f"{boundary}.{variant}.golden_vs_golden", golden_variants[variant], golden_output) + ), + "actual_recompute_vs_actual_output": asdict( + stats(f"{boundary}.{variant}.actual_vs_actual", actual_variants[variant], actual_output) + ), + "actual_recompute_vs_golden_output": asdict( + stats(f"{boundary}.{variant}.actual_vs_golden", actual_variants[variant], golden_output) + ), + } + ) + payload[boundary] = { + "direct_input_actual_vs_golden": asdict(stats(f"{boundary}.input_actual_vs_golden", actual_input, golden_input)), + "direct_output_actual_vs_golden": asdict( + stats(f"{boundary}.output_actual_vs_golden", actual_output, golden_output) + ), + "variants": reports, + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--golden-stage", help="HF full-stage golden safetensors for --layer-idx") + parser.add_argument("--actual-stage", help="PegaInfer full-stage actual safetensors for --layer-idx") + parser.add_argument("--golden-one-step", help="one-step or rich-trace golden with final_hidden.bf16") + parser.add_argument("--actual-one-step", help="PegaInfer one-step actual with final_hidden.bf16") + parser.add_argument("--model-safetensors", required=True) + parser.add_argument("--layer-idx", type=int) + parser.add_argument("--eps", type=float, default=1e-6) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--json-out", default="") + args = parser.parse_args() + + if args.device.startswith("cuda") and not torch.cuda.is_available(): + raise RuntimeError(f"{args.device} requested but torch.cuda.is_available() is false") + + model_file = Path(args.model_safetensors) + payload: dict[str, object] = { + "model_safetensors": args.model_safetensors, + "layer_idx": args.layer_idx, + "eps": args.eps, + "device": args.device, + } + + if args.golden_stage and args.actual_stage: + if args.layer_idx is None: + raise ValueError("--layer-idx is required with stage dumps") + golden_stage = load_file(args.golden_stage) + actual_stage = load_file(args.actual_stage) + layer = args.layer_idx + weight = checkpoint_vec(model_file, f"body.layers.{layer}.input_layernorm.weight", args.device) + add_variant_reports( + payload, + boundary=f"layer{layer}.input_norm", + golden_input=golden_stage[f"layer{layer}.input_hidden.bf16"].to(args.device), + golden_output=golden_stage[f"layer{layer}.input_norm.bf16"].to(args.device), + actual_input=actual_stage[f"layer{layer}.input_hidden.bf16"].to(args.device), + actual_output=actual_stage[f"layer{layer}.input_norm.bf16"].to(args.device), + weight=weight, + eps=args.eps, + ) + + if args.golden_stage and args.actual_stage and args.golden_one_step and args.actual_one_step: + if args.layer_idx is None: + raise ValueError("--layer-idx is required with final norm inputs") + golden_stage = load_file(args.golden_stage) + actual_stage = load_file(args.actual_stage) + golden_one = load_file(args.golden_one_step) + actual_one = load_file(args.actual_one_step) + layer = args.layer_idx + weight = checkpoint_vec(model_file, "body.norm.weight", args.device) + add_variant_reports( + payload, + boundary="final_norm", + golden_input=golden_stage[f"layer{layer}.output_hidden.bf16"].to(args.device), + golden_output=golden_one["final_hidden.bf16"].to(args.device), + actual_input=actual_stage[f"layer{layer}.output_hidden.bf16"].to(args.device), + actual_output=actual_one["final_hidden.bf16"].to(args.device), + weight=weight, + eps=args.eps, + ) + + if args.json_out: + out = Path(args.json_out) + out.parent.mkdir(parents=True, exist_ok=True) + out.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8") + + for boundary, report in payload.items(): + if not isinstance(report, dict) or "variants" not in report: + continue + direct = report["direct_output_actual_vs_golden"] + print( + f"{boundary}: direct_output_mean={direct['mean_abs']:.8f} " + f"direct_output_cos={direct['cosine']:.9f}" + ) + for variant in sorted( + report["variants"], + key=lambda item: item["actual_recompute_vs_actual_output"]["mean_abs"], + ): + actual_fit = variant["actual_recompute_vs_actual_output"] + golden_fit = variant["golden_recompute_vs_golden_output"] + actual_to_golden = variant["actual_recompute_vs_golden_output"] + print( + f" {variant['variant']:<24} " + f"golden_fit={golden_fit['mean_abs']:.8f} " + f"actual_fit={actual_fit['mean_abs']:.8f} " + f"actual_vs_golden={actual_to_golden['mean_abs']:.8f}" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/compare_higgs_layer_hidden.py b/tools/accuracy/compare_higgs_layer_hidden.py new file mode 100755 index 000000000..3bccef1fa --- /dev/null +++ b/tools/accuracy/compare_higgs_layer_hidden.py @@ -0,0 +1,104 @@ +#!/usr/bin/env python3 +"""Compare Higgs/Qwen3 per-layer hidden snapshots from HF and PegaInfer.""" + +from __future__ import annotations + +import argparse +from dataclasses import dataclass +from pathlib import Path + +import torch +from safetensors.torch import load_file + +NUM_LAYERS = 36 + + +@dataclass +class DriftStats: + name: str + max_abs: float + mean_abs: float + p99_abs: float + rmse: float + cosine: float + + +def cosine(left: torch.Tensor, right: torch.Tensor) -> float: + return float( + torch.nn.functional.cosine_similarity( + left.float().flatten(), right.float().flatten(), dim=0 + ) + ) + + +def quantile(values: torch.Tensor, q: float) -> float: + if values.numel() == 0: + return 0.0 + return float(torch.quantile(values.float(), q)) + + +def stats(name: str, golden: torch.Tensor, actual: torch.Tensor) -> DriftStats: + if tuple(golden.shape) != tuple(actual.shape): + raise ValueError(f"{name} shape mismatch: golden {tuple(golden.shape)} actual {tuple(actual.shape)}") + delta = actual.float().flatten() - golden.float().flatten() + abs_delta = delta.abs() + return DriftStats( + name=name, + max_abs=float(abs_delta.max()), + mean_abs=float(abs_delta.mean()), + p99_abs=quantile(abs_delta, 0.99), + rmse=float(torch.sqrt(torch.mean(delta * delta))), + cosine=cosine(golden, actual), + ) + + +def prompt_exact(golden: dict[str, torch.Tensor], actual: dict[str, torch.Tensor]) -> bool: + for name in ("prompt.input_ids_padded", "prompt.attention_mask", "prompt.lengths"): + if not torch.equal(golden[name], actual[name]): + return False + return True + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--golden", type=Path, required=True) + parser.add_argument("--actual", type=Path, required=True) + parser.add_argument("--mean-alert", type=float, default=0.003) + parser.add_argument("--cosine-alert", type=float, default=0.9998) + args = parser.parse_args() + + golden = load_file(str(args.golden), device="cpu") + actual = load_file(str(args.actual), device="cpu") + print(f"prompt_exact={prompt_exact(golden, actual)}") + print("name max_abs mean_abs p99_abs rmse cosine") + + rows: list[DriftStats] = [] + rows.append(stats("embedding.last_hidden.bf16", golden["embedding.last_hidden.bf16"], actual["embedding.last_hidden.bf16"])) + for layer_idx in range(NUM_LAYERS): + name = f"layer.{layer_idx:02}.last_hidden.bf16" + rows.append(stats(name, golden[name], actual[name])) + rows.append(stats("final_hidden.bf16", golden["final_hidden.bf16"], actual["final_hidden.bf16"])) + + first_mean_alert = None + first_cos_alert = None + for row in rows: + print( + f"{row.name:28} {row.max_abs:8.6f} {row.mean_abs:10.6f} " + f"{row.p99_abs:10.6f} {row.rmse:10.6f} {row.cosine:12.9f}" + ) + if first_mean_alert is None and row.mean_abs > args.mean_alert: + first_mean_alert = row.name + if first_cos_alert is None and row.cosine < args.cosine_alert: + first_cos_alert = row.name + + print("summary:") + print(f" first_mean_abs_gt_{args.mean_alert:.6f}={first_mean_alert or 'none'}") + print(f" first_cosine_lt_{args.cosine_alert:.9f}={first_cos_alert or 'none'}") + worst_mean = max(rows, key=lambda row: row.mean_abs) + worst_cosine = min(rows, key=lambda row: row.cosine) + print(f" worst_mean_abs={worst_mean.name}:{worst_mean.mean_abs:.6f}") + print(f" worst_cosine={worst_cosine.name}:{worst_cosine.cosine:.9f}") + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/compare_higgs_stage_dump.py b/tools/accuracy/compare_higgs_stage_dump.py new file mode 100755 index 000000000..a1954878c --- /dev/null +++ b/tools/accuracy/compare_higgs_stage_dump.py @@ -0,0 +1,137 @@ +#!/usr/bin/env python3 +"""Compare two Higgs diagnostic stage safetensors dumps.""" + +from __future__ import annotations + +import argparse +import json +from dataclasses import asdict, dataclass +from pathlib import Path + +import torch +from safetensors.torch import load_file + +PROMPT_TENSORS = {"prompt.input_ids_padded", "prompt.attention_mask", "prompt.lengths"} +STAGE_SUFFIX_ORDER = [ + "input_hidden", + "input_norm", + "q_proj", + "k_proj", + "v_proj", + "q_norm", + "k_norm", + "q_norm_rope", + "k_norm_rope", + "attn_output", + "o_proj", + "post_attn_norm", + "gate_proj", + "up_proj", + "silu_mul", + "down_proj", + "output_hidden", +] + + +@dataclass +class DriftStats: + name: str + max_abs: float + mean_abs: float + p99_abs: float + rmse: float + cosine: float + + +def cosine(left: torch.Tensor, right: torch.Tensor) -> float: + return float(torch.nn.functional.cosine_similarity(left.float().flatten(), right.float().flatten(), dim=0)) + + +def quantile(values: torch.Tensor, q: float) -> float: + if values.numel() == 0: + return 0.0 + return float(torch.quantile(values.float(), q)) + + +def stats(name: str, golden: torch.Tensor, actual: torch.Tensor) -> DriftStats: + if tuple(golden.shape) != tuple(actual.shape): + raise ValueError(f"{name} shape mismatch: golden {tuple(golden.shape)} actual {tuple(actual.shape)}") + delta = actual.float().flatten() - golden.float().flatten() + abs_delta = delta.abs() + return DriftStats( + name=name, + max_abs=float(abs_delta.max()), + mean_abs=float(abs_delta.mean()), + p99_abs=quantile(abs_delta, 0.99), + rmse=float(torch.sqrt(torch.mean(delta * delta))), + cosine=cosine(golden, actual), + ) + + +def prompt_exact(golden: dict[str, torch.Tensor], actual: dict[str, torch.Tensor]) -> bool: + return all(torch.equal(golden[name], actual[name]) for name in sorted(PROMPT_TENSORS)) + + +def stage_order(layer_idx: int) -> list[str]: + return [f"layer{layer_idx}.{suffix}.bf16" for suffix in STAGE_SUFFIX_ORDER] + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--golden", type=Path, required=True) + parser.add_argument("--actual", type=Path, required=True) + parser.add_argument("--layer-idx", type=int, default=0) + parser.add_argument("--mean-alert", type=float, default=0.003) + parser.add_argument("--cosine-alert", type=float, default=0.9998) + parser.add_argument("--json-out", type=Path, default=None) + args = parser.parse_args() + + golden = load_file(str(args.golden), device="cpu") + actual = load_file(str(args.actual), device="cpu") + common_set = (set(golden) & set(actual)) - PROMPT_TENSORS + common = [name for name in stage_order(args.layer_idx) if name in common_set] + common.extend(sorted(common_set - set(common))) + if not common: + raise RuntimeError("no common non-prompt tensors to compare") + + print(f"prompt_exact={prompt_exact(golden, actual)}") + print("name max_abs mean_abs p99_abs rmse cosine") + rows = [stats(name, golden[name], actual[name]) for name in common] + first_mean_alert = None + first_cos_alert = None + for row in rows: + print( + f"{row.name:28} {row.max_abs:8.6f} {row.mean_abs:10.6f} " + f"{row.p99_abs:10.6f} {row.rmse:10.6f} {row.cosine:12.9f}" + ) + if first_mean_alert is None and row.mean_abs > args.mean_alert: + first_mean_alert = row.name + if first_cos_alert is None and row.cosine < args.cosine_alert: + first_cos_alert = row.name + worst_mean = max(rows, key=lambda row: row.mean_abs) + worst_cosine = min(rows, key=lambda row: row.cosine) + payload = { + "golden": str(args.golden), + "actual": str(args.actual), + "layer_idx": args.layer_idx, + "prompt_exact": prompt_exact(golden, actual), + "compared": len(rows), + "first_mean_alert": first_mean_alert or "none", + "first_cos_alert": first_cos_alert or "none", + "worst_mean_abs": f"{worst_mean.name}:{worst_mean.mean_abs:.6f}", + "worst_cosine": f"{worst_cosine.name}:{worst_cosine.cosine:.9f}", + "rows": [asdict(row) for row in rows], + } + if args.json_out is not None: + args.json_out.parent.mkdir(parents=True, exist_ok=True) + args.json_out.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8") + print("summary:") + print(f" compared={len(rows)}") + print(f" first_mean_abs_gt_{args.mean_alert:.6f}={first_mean_alert or 'none'}") + print(f" first_cosine_lt_{args.cosine_alert:.9f}={first_cos_alert or 'none'}") + print(f" worst_mean_abs={worst_mean.name}:{worst_mean.mean_abs:.6f}") + print(f" worst_cosine={worst_cosine.name}:{worst_cosine.cosine:.9f}") + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/compare_higgs_trace_dump.py b/tools/accuracy/compare_higgs_trace_dump.py new file mode 100644 index 000000000..c25ed3839 --- /dev/null +++ b/tools/accuracy/compare_higgs_trace_dump.py @@ -0,0 +1,447 @@ +#!/usr/bin/env python3 +"""Compare Higgs-Audio rich trace safetensors dumps. + +The trace golden is intentionally wider than the default one-step fixture. This +tool compares the common tensor surface so it can be used while the PegaInfer +actual trace is still growing stage by stage. +""" + +from __future__ import annotations + +import argparse +import json +import re +from dataclasses import asdict, dataclass +from pathlib import Path + +import torch +from safetensors.torch import load_file + +PROMPT_TENSORS = ( + "prompt.input_ids_padded", + "prompt.attention_mask", + "prompt.lengths", +) +NUM_LAYERS = 36 +FINAL_LAYER_IDX = NUM_LAYERS - 1 +FINAL_LAYER_HIDDEN_TRACE_KEYS = { + f"layer.{FINAL_LAYER_IDX:02}.last_hidden.bf16", + f"layer.{FINAL_LAYER_IDX:02}.sequence_hidden.bf16", +} +STAGE_SUFFIX_ALIASES = { + "input_norm": "input_layernorm.output", + "q_proj": "self_attn.q_proj.output", + "k_proj": "self_attn.k_proj.output", + "v_proj": "self_attn.v_proj.output", + "q_norm": "self_attn.q_norm.output", + "k_norm": "self_attn.k_norm.output", + "o_proj": "self_attn.o_proj.output", + "post_attn_norm": "post_attention_layernorm.output", + "gate_proj": "mlp.gate_proj.output", + "up_proj": "mlp.up_proj.output", + "down_proj": "mlp.down_proj.output", +} + + + +STAGE_EXECUTION_ORDER = { + "input_hidden": 0, + "input_norm": 1, + "q_proj": 2, + "k_proj": 3, + "v_proj": 4, + "q_norm": 5, + "k_norm": 6, + "q_norm_rope": 7, + "k_norm_rope": 8, + "attn_output": 9, + "o_proj": 10, + "post_attn_norm": 11, + "gate_proj": 12, + "up_proj": 13, + "silu_mul": 14, + "down_proj": 15, + "output_hidden": 16, +} + + +def stage_actual_order(name: str) -> tuple[int, int, str]: + match = re.fullmatch(r"layer(\d+)\.([A-Za-z0-9_]+)\.bf16", name) + if match is None: + return (10_000, 10_000, name) + return (int(match.group(1)), STAGE_EXECUTION_ORDER.get(match.group(2), 9_999), name) + +def stage_actual_to_trace_name(actual_name: str) -> str | None: + match = re.fullmatch(r"layer(\d+)\.([A-Za-z0-9_]+)\.bf16", actual_name) + if match is None: + return None + layer_idx = int(match.group(1)) + suffix = match.group(2) + if suffix == "input_hidden": + if layer_idx == 0: + return "embedding.sequence_hidden.bf16" + return f"layer.{layer_idx - 1:02}.sequence_hidden.bf16" + if suffix == "output_hidden": + if layer_idx == FINAL_LAYER_IDX: + return None + return f"layer.{layer_idx:02}.sequence_hidden.bf16" + trace_suffix = STAGE_SUFFIX_ALIASES.get(suffix) + if trace_suffix is None: + return None + return f"layer.{layer_idx:02}.{trace_suffix}.bf16" + + + +@dataclass +class TensorStats: + name: str + dtype: str + shape: list[int] + max_abs: float | None + mean_abs: float | None + p99_abs: float | None + rmse: float | None + cosine: float | None + exact: bool + alert: bool + reason: str + + +@dataclass +class CompareItem: + display_name: str + golden_name: str + actual_name: str + + +def natural_key(name: str) -> tuple: + parts: list[object] = [] + for piece in re.split(r"(\d+)", name): + if piece.isdigit(): + parts.append(int(piece)) + else: + parts.append(piece) + return tuple(parts) + + +def trace_order(name: str) -> tuple: + if name.startswith("prompt."): + group = 0 + elif name.startswith("embedding."): + group = 1 + elif name.startswith("layer."): + group = 2 + elif name.startswith("final_hidden."): + group = 3 + elif name.startswith("audio_head."): + group = 4 + elif name.startswith("audio_"): + group = 5 + else: + group = 9 + return (group, natural_key(name)) + + +def quantile(values: torch.Tensor, q: float) -> float: + if values.numel() == 0: + return 0.0 + return float(torch.quantile(values.float(), q)) + + +def cosine(left: torch.Tensor, right: torch.Tensor) -> float: + left_flat = left.float().flatten() + right_flat = right.float().flatten() + if left_flat.numel() == 0: + return 1.0 + if float(torch.linalg.vector_norm(left_flat)) == 0.0 and float(torch.linalg.vector_norm(right_flat)) == 0.0: + return 1.0 + return float(torch.nn.functional.cosine_similarity(left_flat, right_flat, dim=0)) + + +def last_token_from_sequence(sequence: torch.Tensor, prompt_lengths: torch.Tensor) -> torch.Tensor: + if sequence.ndim < 3: + return sequence + rows = [] + for batch_idx, prompt_len in enumerate(prompt_lengths.tolist()): + rows.append(sequence[batch_idx, int(prompt_len) - 1]) + return torch.stack(rows, dim=0) + + +def align_golden_to_actual( + golden: torch.Tensor, + actual: torch.Tensor, + prompt_lengths: torch.Tensor | None, +) -> torch.Tensor: + if tuple(golden.shape) == tuple(actual.shape): + return golden + if prompt_lengths is not None and golden.ndim >= 3: + last = last_token_from_sequence(golden, prompt_lengths) + if tuple(last.shape) == tuple(actual.shape): + return last + if actual.ndim == 2 and last.ndim > 2 and last.shape[0] == actual.shape[0]: + flattened = last.reshape(last.shape[0], -1) + if tuple(flattened.shape) == tuple(actual.shape): + return flattened + return golden + + +def compare_tensor( + name: str, + golden: torch.Tensor, + actual: torch.Tensor, + *, + prompt_lengths: torch.Tensor | None, + mean_alert: float, + cosine_alert: float, + max_alert: float | None, +) -> TensorStats: + golden = align_golden_to_actual(golden, actual, prompt_lengths) + if tuple(golden.shape) != tuple(actual.shape): + return TensorStats( + name=name, + dtype=f"{golden.dtype}/{actual.dtype}", + shape=list(golden.shape), + max_abs=None, + mean_abs=None, + p99_abs=None, + rmse=None, + cosine=None, + exact=False, + alert=True, + reason=f"shape_mismatch golden={tuple(golden.shape)} actual={tuple(actual.shape)}", + ) + + exact = bool(torch.equal(golden, actual)) + if not (torch.is_floating_point(golden) or torch.is_floating_point(actual)): + return TensorStats( + name=name, + dtype=str(golden.dtype), + shape=list(golden.shape), + max_abs=None, + mean_abs=None, + p99_abs=None, + rmse=None, + cosine=None, + exact=exact, + alert=not exact, + reason="ok" if exact else "integer_mismatch", + ) + + delta = actual.float().flatten() - golden.float().flatten() + abs_delta = delta.abs() + max_abs = float(abs_delta.max()) if abs_delta.numel() else 0.0 + mean_abs = float(abs_delta.mean()) if abs_delta.numel() else 0.0 + p99_abs = quantile(abs_delta, 0.99) + rmse = float(torch.sqrt(torch.mean(delta * delta))) if delta.numel() else 0.0 + cos = cosine(golden, actual) + reasons = [] + if mean_abs > mean_alert: + reasons.append(f"mean_abs>{mean_alert}") + if cos < cosine_alert: + reasons.append(f"cosine<{cosine_alert}") + if max_alert is not None and max_abs > max_alert: + reasons.append(f"max_abs>{max_alert}") + return TensorStats( + name=name, + dtype=str(golden.dtype), + shape=list(golden.shape), + max_abs=max_abs, + mean_abs=mean_abs, + p99_abs=p99_abs, + rmse=rmse, + cosine=cos, + exact=exact, + alert=bool(reasons), + reason=";".join(reasons) if reasons else "ok", + ) + + +def prompt_exact(golden: dict[str, torch.Tensor], actual: dict[str, torch.Tensor]) -> bool: + for name in PROMPT_TENSORS: + if name not in golden or name not in actual: + return False + if not torch.equal(golden[name], actual[name]): + return False + return True + + +def select_names( + golden: dict[str, torch.Tensor], + actual: dict[str, torch.Tensor], + include_regex: str, +) -> tuple[list[str], list[str], list[str]]: + pattern = re.compile(include_regex) if include_regex else None + golden_names = set(golden) - FINAL_LAYER_HIDDEN_TRACE_KEYS + actual_names = set(actual) - FINAL_LAYER_HIDDEN_TRACE_KEYS + if pattern is not None: + golden_names = {name for name in golden_names if pattern.search(name)} + actual_names = {name for name in actual_names if pattern.search(name)} + common = sorted(golden_names & actual_names, key=trace_order) + missing_from_actual = sorted(golden_names - actual_names, key=trace_order) + extra_actual = sorted(actual_names - golden_names, key=trace_order) + return common, missing_from_actual, extra_actual + + +def select_items( + golden: dict[str, torch.Tensor], + actual: dict[str, torch.Tensor], + include_regex: str, + alias_set: str, +) -> tuple[list[CompareItem], list[str], list[str]]: + if alias_set == "none": + common, missing_from_actual, extra_actual = select_names(golden, actual, include_regex) + return [CompareItem(name, name, name) for name in common], missing_from_actual, extra_actual + if alias_set not in {"layer0-stage", "layer-stage"}: + raise ValueError(f"unknown alias set: {alias_set}") + + pattern = re.compile(include_regex) if include_regex else None + items = [] + missing_from_golden = [] + seen_golden = set() + actual_names = sorted(actual, key=stage_actual_order) + for actual_name in actual_names: + golden_name = stage_actual_to_trace_name(actual_name) + if golden_name is None: + continue + if alias_set == "layer0-stage" and not actual_name.startswith("layer0."): + continue + display_name = f"{actual_name} -> {golden_name}" + if pattern is not None and not (pattern.search(actual_name) or pattern.search(golden_name)): + continue + if golden_name not in golden: + missing_from_golden.append(golden_name) + continue + seen_golden.add(golden_name) + items.append(CompareItem(display_name, golden_name, actual_name)) + missing_from_actual = [] + if alias_set == "layer0-stage": + for actual_name in ( + "layer0.input_hidden.bf16", + "layer0.input_norm.bf16", + "layer0.q_proj.bf16", + "layer0.k_proj.bf16", + "layer0.v_proj.bf16", + "layer0.q_norm.bf16", + "layer0.k_norm.bf16", + "layer0.o_proj.bf16", + "layer0.post_attn_norm.bf16", + "layer0.gate_proj.bf16", + "layer0.up_proj.bf16", + "layer0.down_proj.bf16", + "layer0.output_hidden.bf16", + ): + if actual_name not in actual: + missing_from_actual.append(actual_name) + return items, missing_from_actual, sorted(set(missing_from_golden), key=trace_order) + + +def format_float(value: float | None) -> str: + if value is None: + return " n/a" + return f"{value:10.6f}" + + +def write_json(path: Path, payload: dict[str, object]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n") + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--golden", type=Path, required=True) + parser.add_argument("--actual", type=Path, required=True) + parser.add_argument("--mean-alert", type=float, default=0.003) + parser.add_argument("--cosine-alert", type=float, default=0.9998) + parser.add_argument("--max-alert", type=float, default=None) + parser.add_argument("--include-regex", default="") + parser.add_argument( + "--alias-set", + choices=("none", "layer0-stage", "layer-stage"), + default="none", + help="Optional built-in mapping from a partial actual dump schema to the trace golden schema.", + ) + parser.add_argument("--top", type=int, default=120) + parser.add_argument("--json-out", type=Path, default=None) + parser.add_argument("--require-all", action="store_true") + parser.add_argument("--fail-on-alert", action="store_true") + args = parser.parse_args() + + golden = load_file(str(args.golden), device="cpu") + actual = load_file(str(args.actual), device="cpu") + items, missing_from_actual, extra_actual = select_items( + golden, actual, args.include_regex, args.alias_set + ) + if not items: + raise RuntimeError("no comparable tensors to compare") + prompt_lengths = golden.get("prompt.lengths") + + rows = [ + compare_tensor( + item.display_name, + golden[item.golden_name], + actual[item.actual_name], + prompt_lengths=prompt_lengths, + mean_alert=args.mean_alert, + cosine_alert=args.cosine_alert, + max_alert=args.max_alert, + ) + for item in items + ] + alerts = [row for row in rows if row.alert] + first_alert = alerts[0].name if alerts else "none" + floating_rows = [row for row in rows if row.mean_abs is not None and row.cosine is not None] + worst_mean = max(floating_rows, key=lambda row: row.mean_abs or 0.0, default=None) + worst_cosine = min(floating_rows, key=lambda row: row.cosine if row.cosine is not None else 1.0, default=None) + + print(f"prompt_exact={prompt_exact(golden, actual)}") + print(f"common_tensors={len(items)}") + print(f"missing_from_actual={len(missing_from_actual)}") + print(f"extra_actual={len(extra_actual)}") + print("name max_abs mean_abs p99_abs rmse cosine exact alert reason") + for row in rows[: args.top]: + print( + f"{row.name:58} {format_float(row.max_abs)} {format_float(row.mean_abs)} " + f"{format_float(row.p99_abs)} {format_float(row.rmse)} {format_float(row.cosine)} " + f"{str(row.exact):>5} {str(row.alert):>6} {row.reason}" + ) + if len(rows) > args.top: + print(f"... truncated {len(rows) - args.top} row(s); use --top to print more") + print("summary:") + print(f" compared={len(rows)}") + print(f" alerts={len(alerts)}") + print(f" first_alert={first_alert}") + print(f" worst_mean_abs={(worst_mean.name + ':' + f'{worst_mean.mean_abs:.6f}') if worst_mean else 'none'}") + print(f" worst_cosine={(worst_cosine.name + ':' + f'{worst_cosine.cosine:.9f}') if worst_cosine else 'none'}") + if missing_from_actual[:10]: + print(f" missing_from_actual_first10={missing_from_actual[:10]}") + if extra_actual[:10]: + print(f" extra_actual_first10={extra_actual[:10]}") + + if args.json_out is not None: + write_json( + args.json_out, + { + "golden": str(args.golden), + "actual": str(args.actual), + "prompt_exact": prompt_exact(golden, actual), + "common_tensors": len(items), + "alias_set": args.alias_set, + "items": [asdict(item) for item in items], + "missing_from_actual": missing_from_actual, + "extra_actual": extra_actual, + "alerts": len(alerts), + "first_alert": first_alert, + "worst_mean_abs": asdict(worst_mean) if worst_mean else None, + "worst_cosine": asdict(worst_cosine) if worst_cosine else None, + "rows": [asdict(row) for row in rows], + }, + ) + + if args.require_all and missing_from_actual: + raise SystemExit(2) + if args.fail_on_alert and alerts: + raise SystemExit(1) + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/dump_higgs_layer0_stages_golden.py b/tools/accuracy/dump_higgs_layer0_stages_golden.py new file mode 100755 index 000000000..0fb089864 --- /dev/null +++ b/tools/accuracy/dump_higgs_layer0_stages_golden.py @@ -0,0 +1,239 @@ +#!/usr/bin/env python3 +"""Dump HF Higgs/Qwen3 prefill stage snapshots for the one-step prompt. + +The filename is kept for compatibility with earlier layer-0 workflows, but the +tool supports any decoder layer via --layer-idx. +""" + +from __future__ import annotations + +import argparse +import json +import platform +from pathlib import Path + +import torch +from safetensors.torch import save_file + +from transformers.models.qwen3.modeling_qwen3 import apply_rotary_pos_emb + +from dump_higgs_one_step_golden import ( + DEFAULT_MODEL_ID, + DEFAULT_PROMPTS, + DEFAULT_REVISION, + HiggsTokenizerAdapter, + load_backbone, + load_sglang_omni_reference, + load_tokenizer, + sha256_file, +) + + +def last_token(hidden: torch.Tensor, row_idx: torch.Tensor, prompt_lens: torch.Tensor) -> torch.Tensor: + return hidden[row_idx, prompt_lens - 1, :].detach().contiguous().clone() + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--model-id", default=DEFAULT_MODEL_ID) + parser.add_argument("--revision", default=DEFAULT_REVISION) + parser.add_argument("--snapshot-dir", required=True) + parser.add_argument("--out", required=True) + parser.add_argument("--layer-idx", type=int, default=0) + parser.add_argument("--prompt", action="append", default=[]) + parser.add_argument("--device", default="cuda:0") + parser.add_argument("--sglang-omni-src", default="") + parser.add_argument("--require-sglang-omni-source", action="store_true") + args = parser.parse_args() + + torch.set_grad_enabled(False) + torch.manual_seed(0) + if args.device.startswith("cuda") and not torch.cuda.is_available(): + raise RuntimeError("CUDA requested but torch.cuda.is_available() is false") + + snapshot_dir = Path(args.snapshot_dir) + config_path = snapshot_dir / "config.json" + tokenizer_path = snapshot_dir / "tokenizer.json" + model_file = snapshot_dir / "model.safetensors" + index_path = snapshot_dir / "model.safetensors.index.json" + for path in (config_path, tokenizer_path, model_file, index_path): + if not path.exists(): + raise FileNotFoundError(path) + + config = json.load(open(config_path)) + text_cfg = dict(config["text_config"]) + prompts = args.prompt or list(DEFAULT_PROMPTS) + reference = load_sglang_omni_reference( + args.sglang_omni_src, + required=args.require_sglang_omni_source, + ) + adapter_cls = reference.get("tokenizer_adapter_cls", HiggsTokenizerAdapter) + adapter = load_tokenizer(snapshot_dir, adapter_cls) + prompt_ids = [adapter.build_prompt(prompt) for prompt in prompts] + max_len = max(len(ids) for ids in prompt_ids) + pad_id = int(text_cfg.get("eos_token_id") or 151643) + input_ids = torch.full((len(prompt_ids), max_len), pad_id, dtype=torch.long, device=args.device) + attention_mask = torch.zeros((len(prompt_ids), max_len), dtype=torch.long, device=args.device) + for row, ids in enumerate(prompt_ids): + input_ids[row, : len(ids)] = torch.tensor(ids, dtype=torch.long, device=args.device) + attention_mask[row, : len(ids)] = 1 + prompt_lens = torch.tensor([len(ids) for ids in prompt_ids], dtype=torch.int64) + prompt_lens_device = prompt_lens.to(args.device) + row_idx = torch.arange(len(prompt_ids), device=args.device) + + backbone = load_backbone(model_file, text_cfg, args.device) + if args.layer_idx < 0 or args.layer_idx >= len(backbone.layers): + raise ValueError(f"--layer-idx {args.layer_idx} out of range for {len(backbone.layers)} layers") + layer = backbone.layers[args.layer_idx] + stage_prefix = f"layer{args.layer_idx}" + stages: dict[str, torch.Tensor] = {} + q_norm_all: torch.Tensor | None = None + k_norm_all: torch.Tensor | None = None + + def capture(name: str): + def hook(_module, _inputs, output): + tensor = output[0] if isinstance(output, tuple) else output + stages[name] = last_token(tensor, row_idx, prompt_lens_device) + + return hook + + def capture_pre(name: str): + def hook(_module, inputs): + stages[name] = last_token(inputs[0], row_idx, prompt_lens_device) + + return hook + + def capture_q_norm(_module, _inputs, output): + nonlocal q_norm_all + q_norm_all = output.detach().contiguous().clone() + + def capture_k_norm(_module, _inputs, output): + nonlocal k_norm_all + k_norm_all = output.detach().contiguous().clone() + + hooks = [ + layer.register_forward_pre_hook(capture_pre(f"{stage_prefix}.input_hidden.bf16")), + layer.input_layernorm.register_forward_hook(capture(f"{stage_prefix}.input_norm.bf16")), + layer.self_attn.q_proj.register_forward_hook(capture(f"{stage_prefix}.q_proj.bf16")), + layer.self_attn.k_proj.register_forward_hook(capture(f"{stage_prefix}.k_proj.bf16")), + layer.self_attn.v_proj.register_forward_hook(capture(f"{stage_prefix}.v_proj.bf16")), + layer.self_attn.q_norm.register_forward_hook(capture_q_norm), + layer.self_attn.k_norm.register_forward_hook(capture_k_norm), + layer.self_attn.o_proj.register_forward_pre_hook(capture_pre(f"{stage_prefix}.attn_output.bf16")), + layer.self_attn.o_proj.register_forward_hook(capture(f"{stage_prefix}.o_proj.bf16")), + layer.post_attention_layernorm.register_forward_hook(capture(f"{stage_prefix}.post_attn_norm.bf16")), + layer.mlp.gate_proj.register_forward_hook(capture(f"{stage_prefix}.gate_proj.bf16")), + layer.mlp.up_proj.register_forward_hook(capture(f"{stage_prefix}.up_proj.bf16")), + layer.mlp.down_proj.register_forward_pre_hook(capture_pre(f"{stage_prefix}.silu_mul.bf16")), + layer.mlp.down_proj.register_forward_hook(capture(f"{stage_prefix}.down_proj.bf16")), + layer.register_forward_hook(capture(f"{stage_prefix}.output_hidden.bf16")), + ] + try: + with torch.inference_mode(): + backbone( + input_ids=input_ids, + attention_mask=attention_mask, + use_cache=False, + return_dict=True, + ) + finally: + for hook in hooks: + hook.remove() + + if q_norm_all is None or k_norm_all is None: + raise RuntimeError("missing q/k norm snapshots") + position_ids = torch.arange(max_len, device=args.device).unsqueeze(0) + + def norm_to_bhsd(name: str, tensor: torch.Tensor) -> torch.Tensor: + if tensor.ndim != 4: + raise RuntimeError(f"{name} expected rank-4 q/k norm output, got {tuple(tensor.shape)}") + if tensor.shape[1] == max_len: + return tensor.transpose(1, 2).contiguous() + if tensor.shape[2] == max_len: + return tensor.contiguous() + raise RuntimeError(f"{name} cannot infer seq axis from shape {tuple(tensor.shape)}") + + q_norm_bhsd = norm_to_bhsd("q_norm", q_norm_all) + k_norm_bhsd = norm_to_bhsd("k_norm", k_norm_all) + cos, sin = backbone.rotary_emb(q_norm_bhsd, position_ids) + stages[f"{stage_prefix}.q_norm.bf16"] = last_token( + q_norm_bhsd.transpose(1, 2).reshape(len(prompt_ids), max_len, -1), + row_idx, + prompt_lens_device, + ) + stages[f"{stage_prefix}.k_norm.bf16"] = last_token( + k_norm_bhsd.transpose(1, 2).reshape(len(prompt_ids), max_len, -1), + row_idx, + prompt_lens_device, + ) + q_rope, k_rope = apply_rotary_pos_emb(q_norm_bhsd, k_norm_bhsd, cos, sin) + stages[f"{stage_prefix}.q_norm_rope.bf16"] = last_token( + q_rope.transpose(1, 2).reshape(len(prompt_ids), max_len, -1), + row_idx, + prompt_lens_device, + ) + stages[f"{stage_prefix}.k_norm_rope.bf16"] = last_token( + k_rope.transpose(1, 2).reshape(len(prompt_ids), max_len, -1), + row_idx, + prompt_lens_device, + ) + + expected = [ + f"{stage_prefix}.input_hidden.bf16", + f"{stage_prefix}.input_norm.bf16", + f"{stage_prefix}.q_proj.bf16", + f"{stage_prefix}.k_proj.bf16", + f"{stage_prefix}.v_proj.bf16", + f"{stage_prefix}.q_norm.bf16", + f"{stage_prefix}.k_norm.bf16", + f"{stage_prefix}.q_norm_rope.bf16", + f"{stage_prefix}.k_norm_rope.bf16", + f"{stage_prefix}.attn_output.bf16", + f"{stage_prefix}.o_proj.bf16", + f"{stage_prefix}.post_attn_norm.bf16", + f"{stage_prefix}.gate_proj.bf16", + f"{stage_prefix}.up_proj.bf16", + f"{stage_prefix}.silu_mul.bf16", + f"{stage_prefix}.down_proj.bf16", + f"{stage_prefix}.output_hidden.bf16", + ] + missing = [name for name in expected if name not in stages] + if missing: + raise RuntimeError(f"missing stage snapshots: {missing}") + + tensors = { + "prompt.input_ids_padded": input_ids.cpu().to(torch.int64), + "prompt.attention_mask": attention_mask.cpu().to(torch.int64), + "prompt.lengths": prompt_lens.cpu(), + } + for name in expected: + tensors[name] = stages[name].cpu().to(torch.bfloat16) + + metadata = { + "fixture_kind": "higgs-layer-stage-golden", + "schema_version": "1", + "layer_idx": str(args.layer_idx), + "model_id": args.model_id, + "model_revision": args.revision, + "reference": "Transformers Qwen3Model per-layer module hooks", + "sglang_omni_source_dir": reference.get("source_dir", ""), + "sglang_omni_source_commit": reference.get("source_commit", ""), + "prompt_count": str(len(prompts)), + "prompts_json": json.dumps(prompts, ensure_ascii=False), + "config_sha256": sha256_file(config_path), + "tokenizer_json_sha256": sha256_file(tokenizer_path), + "model_index_sha256": sha256_file(index_path), + "python": platform.python_version(), + "torch": torch.__version__, + "transformers": __import__("transformers").__version__, + "device": torch.cuda.get_device_name(0) if args.device.startswith("cuda") else args.device, + } + out_path = Path(args.out) + out_path.parent.mkdir(parents=True, exist_ok=True) + save_file(tensors, str(out_path), metadata=metadata) + print(f"wrote {out_path} size={out_path.stat().st_size}") + print(f"stages {len(expected)}") + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/dump_higgs_layer_hidden_golden.py b/tools/accuracy/dump_higgs_layer_hidden_golden.py new file mode 100644 index 000000000..916f932de --- /dev/null +++ b/tools/accuracy/dump_higgs_layer_hidden_golden.py @@ -0,0 +1,146 @@ +#!/usr/bin/env python3 +"""Dump HF Higgs/Qwen3 per-layer hidden snapshots for the one-step prompt.""" + +from __future__ import annotations + +import argparse +import json +import platform +from pathlib import Path + +import torch +from safetensors.torch import save_file + +from dump_higgs_one_step_golden import ( + DEFAULT_MODEL_ID, + DEFAULT_PROMPTS, + DEFAULT_REVISION, + load_backbone, + load_tokenizer, + sha256_file, +) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--model-id", default=DEFAULT_MODEL_ID) + parser.add_argument("--revision", default=DEFAULT_REVISION) + parser.add_argument("--snapshot-dir", required=True) + parser.add_argument("--out", required=True) + parser.add_argument("--prompt", action="append", default=[]) + parser.add_argument("--device", default="cuda:0") + args = parser.parse_args() + + torch.set_grad_enabled(False) + torch.manual_seed(0) + if args.device.startswith("cuda") and not torch.cuda.is_available(): + raise RuntimeError("CUDA requested but torch.cuda.is_available() is false") + + snapshot_dir = Path(args.snapshot_dir) + config_path = snapshot_dir / "config.json" + tokenizer_path = snapshot_dir / "tokenizer.json" + model_file = snapshot_dir / "model.safetensors" + index_path = snapshot_dir / "model.safetensors.index.json" + for path in (config_path, tokenizer_path, model_file, index_path): + if not path.exists(): + raise FileNotFoundError(path) + + config = json.load(open(config_path)) + text_cfg = dict(config["text_config"]) + prompts = args.prompt or list(DEFAULT_PROMPTS) + adapter = load_tokenizer(snapshot_dir) + prompt_ids = [adapter.build_prompt(prompt) for prompt in prompts] + max_len = max(len(ids) for ids in prompt_ids) + pad_id = int(text_cfg.get("eos_token_id") or 151643) + input_ids = torch.full( + (len(prompt_ids), max_len), pad_id, dtype=torch.long, device=args.device + ) + attention_mask = torch.zeros( + (len(prompt_ids), max_len), dtype=torch.long, device=args.device + ) + for row, ids in enumerate(prompt_ids): + input_ids[row, : len(ids)] = torch.tensor( + ids, dtype=torch.long, device=args.device + ) + attention_mask[row, : len(ids)] = 1 + prompt_lens = torch.tensor([len(ids) for ids in prompt_ids], dtype=torch.int64) + prompt_lens_device = prompt_lens.to(args.device) + row_idx = torch.arange(len(prompt_ids), device=args.device) + + backbone = load_backbone(model_file, text_cfg, args.device) + embedding_hidden: torch.Tensor | None = None + layer_snapshots: list[torch.Tensor | None] = [None] * len(backbone.layers) + + def make_hook(layer_idx: int): + def hook(_module, _inputs, output): + hidden = output[0] if isinstance(output, tuple) else output + layer_snapshots[layer_idx] = ( + hidden[row_idx, prompt_lens_device - 1, :].detach().contiguous().clone() + ) + + return hook + + def embedding_hook(_module, _inputs, output): + nonlocal embedding_hidden + embedding_hidden = output[row_idx, prompt_lens_device - 1, :].detach().contiguous().clone() + + hooks = [backbone.embed_tokens.register_forward_hook(embedding_hook)] + hooks.extend(layer.register_forward_hook(make_hook(i)) for i, layer in enumerate(backbone.layers)) + try: + with torch.inference_mode(): + out = backbone( + input_ids=input_ids, + attention_mask=attention_mask, + use_cache=False, + return_dict=True, + ) + finally: + for hook in hooks: + hook.remove() + + final_hidden = out.last_hidden_state[ + row_idx, prompt_lens_device - 1, : + ].detach().contiguous() + if embedding_hidden is None: + raise RuntimeError("missing embedding snapshot") + tensors = { + "prompt.input_ids_padded": input_ids.cpu().to(torch.int64), + "prompt.attention_mask": attention_mask.cpu().to(torch.int64), + "prompt.lengths": prompt_lens.cpu(), + "embedding.last_hidden.bf16": embedding_hidden.cpu().to(torch.bfloat16), + "final_hidden.bf16": final_hidden.cpu().to(torch.bfloat16), + } + for layer_idx, hidden in enumerate(layer_snapshots): + if hidden is None: + raise RuntimeError(f"missing layer snapshot {layer_idx}") + tensors[f"layer.{layer_idx:02}.last_hidden.bf16"] = hidden.cpu().to(torch.bfloat16) + + metadata = { + "fixture_kind": "higgs-prefill-layer-hidden-golden", + "schema_version": "1", + "model_id": args.model_id, + "model_revision": args.revision, + "reference": "Transformers Qwen3Model forward hooks after each decoder layer plus final model norm", + "prompt_count": str(len(prompts)), + "prompts_json": json.dumps(prompts, ensure_ascii=False), + "hidden_size": str(text_cfg["hidden_size"]), + "num_hidden_layers": str(text_cfg["num_hidden_layers"]), + "config_sha256": sha256_file(config_path), + "tokenizer_json_sha256": sha256_file(tokenizer_path), + "model_index_sha256": sha256_file(index_path), + "python": platform.python_version(), + "torch": torch.__version__, + "transformers": __import__("transformers").__version__, + "device": torch.cuda.get_device_name(0) + if args.device.startswith("cuda") + else args.device, + } + out_path = Path(args.out) + out_path.parent.mkdir(parents=True, exist_ok=True) + save_file(tensors, str(out_path), metadata=metadata) + print(f"wrote {out_path} size={out_path.stat().st_size}") + print(f"layers {len(layer_snapshots)}") + + +if __name__ == "__main__": + main() diff --git a/tools/accuracy/dump_higgs_one_step_golden.py b/tools/accuracy/dump_higgs_one_step_golden.py new file mode 100644 index 000000000..c58a2753c --- /dev/null +++ b/tools/accuracy/dump_higgs_one_step_golden.py @@ -0,0 +1,499 @@ +#!/usr/bin/env python3 +"""Dump a Higgs Audio one-step audio-logits golden. + +This generator intentionally uses the SGLang-Omni Higgs prompt/head contract +without importing the full SGLang server stack. The prompt builder mirrors +`sglang_omni.models.higgs_tts.text_tokenizer.HiggsTokenizerAdapter`; the fused +audio head mirrors `HiggsFusedMultiTextHead.generate`. +""" + +from __future__ import annotations + +import argparse +import hashlib +import importlib +import json +import platform +import subprocess +import sys +from pathlib import Path +from typing import Any + +import torch +import torch.nn.functional as F +from huggingface_hub import hf_hub_download, snapshot_download +from safetensors import safe_open +from safetensors.torch import save_file +from tokenizers import Tokenizer +from transformers import PreTrainedTokenizerFast +from transformers.models.qwen3.configuration_qwen3 import Qwen3Config +from transformers.models.qwen3.modeling_qwen3 import Qwen3Model, Qwen3RotaryEmbedding + +AUDIO_PLACEHOLDER_ID = -100 +REQUIRED_SPECIALS = ("<|tts|>", "<|ref_audio|>", "<|text|>", "<|audio|>") +DEFAULT_MODEL_ID = "bosonai/higgs-tts-3-4b" +DEFAULT_REVISION = "7556c17e05201fccd9c8cc120bc216dcc7b5d561" +DEFAULT_PROMPTS = ("Hello from PegaInfer.",) +TRACE_MODULE_SUFFIXES = ( + "input_layernorm", + "self_attn.q_proj", + "self_attn.k_proj", + "self_attn.v_proj", + "self_attn.q_norm", + "self_attn.k_norm", + "self_attn.o_proj", + "post_attention_layernorm", + "mlp.gate_proj", + "mlp.up_proj", + "mlp.down_proj", +) + + +def git_short_commit(path: Path) -> str: + result = subprocess.run( + ["git", "-C", str(path), "rev-parse", "--short", "HEAD"], + check=False, + capture_output=True, + text=True, + ) + return result.stdout.strip() if result.returncode == 0 else "unknown" + + +def load_sglang_omni_reference( + source_dir: str, + *, + required: bool, +) -> dict[str, Any]: + if not source_dir: + if required: + raise ValueError("--require-sglang-omni-source requires --sglang-omni-src") + return {} + + src = Path(source_dir).resolve() + if not (src / "sglang_omni/models/higgs_tts/text_tokenizer.py").exists(): + raise FileNotFoundError(f"SGLang-Omni source missing Higgs tokenizer: {src}") + sys.path.insert(0, str(src)) + try: + tokenizer_mod = importlib.import_module( + "sglang_omni.models.higgs_tts.text_tokenizer" + ) + modeling_mod = importlib.import_module("sglang_omni.models.higgs_tts.modeling") + except Exception: + if required: + raise + return {} + + return { + "source_dir": str(src), + "source_commit": git_short_commit(src), + "tokenizer_adapter_cls": tokenizer_mod.HiggsTokenizerAdapter, + "fused_head_cls": modeling_mod.HiggsFusedMultiTextHead, + } + + +class HiggsTokenizerAdapter: + def __init__(self, tokenizer: Any) -> None: + self._tok = tokenizer + vocab = dict(tokenizer.get_added_vocab()) + missing = [t for t in REQUIRED_SPECIALS if t not in vocab] + if missing: + raise ValueError(f"Tokenizer is missing Higgs TTS specials: {missing}") + self.tts_id = int(vocab["<|tts|>"]) + self.ref_audio_id = int(vocab["<|ref_audio|>"]) + self.text_id = int(vocab["<|text|>"]) + self.audio_id = int(vocab["<|audio|>"]) + self.ref_text_id = vocab.get("<|ref_text|>") + + def build_prompt( + self, + prompt_text: str, + *, + num_ref_tokens: int = 0, + reference_text: str | None = None, + ) -> list[int]: + if num_ref_tokens < 0: + raise ValueError(f"num_ref_tokens must be >= 0, got {num_ref_tokens}") + ids: list[int] = [self.tts_id] + if reference_text and num_ref_tokens > 0 and self.ref_text_id is not None: + ids.append(int(self.ref_text_id)) + ids.extend(self._tok.encode(reference_text, add_special_tokens=False)) + if num_ref_tokens > 0: + ids.append(self.ref_audio_id) + ids.extend([AUDIO_PLACEHOLDER_ID] * num_ref_tokens) + ids.append(self.text_id) + ids.extend(self._tok.encode(prompt_text, add_special_tokens=False)) + ids.append(self.audio_id) + return [int(x) for x in ids] + + +def sha256_file(path: str | Path, chunk_size: int = 1024 * 1024) -> str: + h = hashlib.sha256() + with open(path, "rb") as f: + while True: + b = f.read(chunk_size) + if not b: + break + h.update(b) + return h.hexdigest() + + +def remap_backbone_key(src: str) -> str | None: + if src == "tied.embedding.text_embedding.weight": + return "embed_tokens.weight" + if src.startswith("body.layers."): + return src.removeprefix("body.") + if src.startswith("body.norm."): + return src.removeprefix("body.") + return None + + +def load_backbone(model_file: Path, text_cfg: dict[str, Any], device: str) -> Qwen3Model: + cfg = Qwen3Config(**text_cfg) + cfg._attn_implementation = "sdpa" + with torch.device("meta"): + model = Qwen3Model(cfg) + model.to_empty(device=device) + # to_empty() does not populate non-persistent RoPE buffers created on the + # meta device, so rebuild rotary_emb on the real device before loading params. + model.rotary_emb = Qwen3RotaryEmbedding(cfg, device=device) + model.to(dtype=torch.bfloat16) + model.eval() + params = dict(model.named_parameters()) + loaded: set[str] = set() + with safe_open(str(model_file), framework="pt", device="cpu") as f: + for src in f.keys(): + dst = remap_backbone_key(src) + if dst is None: + continue + if dst not in params: + raise KeyError(f"remapped key {src} -> {dst}, but Qwen3Model has no such parameter") + p = params[dst] + t = f.get_tensor(src) + if tuple(t.shape) != tuple(p.shape): + raise ValueError(f"shape mismatch {src}->{dst}: ckpt {tuple(t.shape)} vs model {tuple(p.shape)}") + p.data.copy_(t.to(device=device, dtype=p.dtype, non_blocking=False)) + loaded.add(dst) + missing = sorted(set(params) - loaded) + if missing: + raise RuntimeError(f"missing {len(missing)} backbone parameters, first: {missing[:8]}") + return model + + +def load_modality_head_weight(model_file: Path, device: str) -> torch.Tensor: + key = "tied.embedding.modality_embeddings.0.embedding.weight" + with safe_open(str(model_file), framework="pt", device="cpu") as f: + if key not in f.keys(): + raise KeyError(f"missing fused modality embedding/head weight {key}") + weight = f.get_tensor(key) + if tuple(weight.shape) != (8208, 2560): + raise ValueError(f"unexpected modality head weight shape {tuple(weight.shape)}") + return weight.to(device=device, dtype=torch.bfloat16) + + +def load_tokenizer(snapshot_dir: Path, adapter_cls: type[Any]) -> Any: + raw = Tokenizer.from_file(str(snapshot_dir / "tokenizer.json")) + tokenizer = PreTrainedTokenizerFast(tokenizer_object=raw) + return adapter_cls(tokenizer) + + +def trace_tensor(tensor: torch.Tensor) -> torch.Tensor: + if torch.is_floating_point(tensor): + return tensor.detach().cpu().to(torch.bfloat16).contiguous() + return tensor.detach().cpu().contiguous() + + +def first_tensor(value: Any) -> torch.Tensor | None: + if isinstance(value, torch.Tensor): + return value + if isinstance(value, (list, tuple)): + for item in value: + tensor = first_tensor(item) + if tensor is not None: + return tensor + return None + + +def module_trace_name(module_name: str) -> str | None: + if not module_name.startswith("layers."): + return None + parts = module_name.split(".", 2) + if len(parts) != 3 or not parts[1].isdigit(): + return None + layer_idx = int(parts[1]) + suffix = parts[2] + if suffix not in TRACE_MODULE_SUFFIXES: + return None + return f"layer.{layer_idx:02}.{suffix}.output.bf16" + + +def register_trace_hooks(model: Qwen3Model, tensors: dict[str, torch.Tensor]) -> list[Any]: + handles = [] + + def make_hook(trace_name: str): + def hook(_module: Any, _inputs: tuple[Any, ...], output: Any) -> None: + tensor = first_tensor(output) + if tensor is not None: + tensors[trace_name] = trace_tensor(tensor) + + return hook + + for module_name, module in model.named_modules(): + trace_name = module_trace_name(module_name) + if trace_name is not None: + handles.append(module.register_forward_hook(make_hook(trace_name))) + return handles + + +def compute_modality_logits( + last_hidden: torch.Tensor, + modality_weight: torch.Tensor, + audio_cfg: dict[str, Any], + reference: dict[str, Any], +) -> torch.Tensor: + head_cls = reference.get("fused_head_cls") + if head_cls is None: + logits = F.linear(last_hidden, modality_weight) + return logits.reshape( + last_hidden.shape[0], + int(audio_cfg["num_codebooks"]), + int(audio_cfg["vocab_size"]), + ) + + head = head_cls( + num_codebooks=int(audio_cfg["num_codebooks"]), + vocab_size=int(audio_cfg["vocab_size"]), + hidden_size=last_hidden.shape[-1], + ).to(device=last_hidden.device, dtype=torch.bfloat16) + head.eval() + with torch.no_grad(): + head.weight.copy_(modality_weight) + return head.generate(last_hidden) + + +def add_hidden_state_trace( + tensors: dict[str, torch.Tensor], + hidden_states: tuple[torch.Tensor, ...] | None, + prompt_lens: torch.Tensor, +) -> None: + if hidden_states is None: + return + # Transformers Qwen3 returns output_hidden_states as: + # embedding, layer0_output, ..., layer34_output, final_norm_output + # The raw layer35 decoder output is not present in this tuple because the + # final item is appended after model.norm. Store that tensor as final_hidden + # only; labeling it layer.35.* creates a false last-layer divergence. + last_idx = len(hidden_states) - 1 + cpu_lens = prompt_lens.cpu() + for idx, state in enumerate(hidden_states): + if idx == last_idx: + continue + if idx == 0: + sequence_name = "embedding.sequence_hidden.bf16" + last_name = "embedding.last_hidden.bf16" + else: + sequence_name = f"layer.{idx - 1:02}.sequence_hidden.bf16" + last_name = f"layer.{idx - 1:02}.last_hidden.bf16" + tensors[sequence_name] = trace_tensor(state) + rows = [] + for batch_idx, prompt_len in enumerate(cpu_lens.tolist()): + rows.append(state[batch_idx, int(prompt_len) - 1, :]) + tensors[last_name] = trace_tensor(torch.stack(rows, dim=0)) + + +def write_trace_file( + out_path: Path, + *, + base_tensors: dict[str, torch.Tensor], + trace_tensors: dict[str, torch.Tensor], + hidden_states: tuple[torch.Tensor, ...] | None, + prompt_lens: torch.Tensor, + last_hidden: torch.Tensor, + logits: torch.Tensor, + logprobs: torch.Tensor, + metadata: dict[str, str], +) -> int: + tensors = dict(base_tensors) + add_hidden_state_trace(tensors, hidden_states, prompt_lens) + tensors.update(trace_tensors) + tensors.update( + { + "audio_head.input_hidden.bf16": last_hidden.cpu().to(torch.bfloat16), + "audio_head.flat_logits.f32": logits.reshape(logits.shape[0], -1) + .cpu() + .to(torch.float32), + "audio_logprobs.f32": logprobs.cpu().to(torch.float32), + } + ) + trace_metadata = dict(metadata) + trace_metadata.update( + { + "fixture_kind": "higgs-one-step-trace-golden", + "schema_version": "3", + "trace_contract": "prompt;embedding;per-layer hidden except final decoder raw hidden;per-layer module outputs;qkv/mlp stage hooks;audio-head logits/logprobs/topk/argmax", + "trace_tensor_count": str(len(tensors)), + "trace_module_suffixes": ";".join(TRACE_MODULE_SUFFIXES), + } + ) + out_path.parent.mkdir(parents=True, exist_ok=True) + save_file(tensors, str(out_path), metadata=trace_metadata) + return len(tensors) + + +def main() -> None: + ap = argparse.ArgumentParser() + ap.add_argument("--model-id", default=DEFAULT_MODEL_ID) + ap.add_argument("--revision", default=DEFAULT_REVISION) + ap.add_argument("--snapshot-dir", default="") + ap.add_argument("--out", default="test_data/higgs-one-step-audio-logits.safetensors") + ap.add_argument( + "--trace-out", + default="", + help="Optional rich trace safetensors path with prompt, per-layer, per-module, and audio-head intermediate tensors.", + ) + ap.add_argument("--prompt", action="append", default=[]) + ap.add_argument("--device", default="cuda:0") + ap.add_argument("--download", action="store_true") + ap.add_argument( + "--sglang-omni-src", + default="", + help="Optional SGLang-Omni source tree; imports Higgs tokenizer/head modules directly.", + ) + ap.add_argument( + "--require-sglang-omni-source", + action="store_true", + help="Fail if --sglang-omni-src cannot provide the Higgs reference modules.", + ) + args = ap.parse_args() + + torch.set_grad_enabled(False) + torch.manual_seed(0) + if args.device.startswith("cuda") and not torch.cuda.is_available(): + raise RuntimeError("CUDA requested but torch.cuda.is_available() is false") + + if args.snapshot_dir: + snapshot_dir = Path(args.snapshot_dir) + else: + local_dir = f"models/higgs-tts-3-4b-{args.revision}" + if args.download or not Path(local_dir, "model.safetensors").exists(): + snapshot_download( + repo_id=args.model_id, + revision=args.revision, + local_dir=local_dir, + local_dir_use_symlinks=False, + resume_download=True, + ) + else: + for name in ("config.json", "tokenizer.json", "tokenizer_config.json", "model.safetensors.index.json"): + hf_hub_download(args.model_id, name, revision=args.revision) + snapshot_dir = Path(local_dir) + + config_path = snapshot_dir / "config.json" + tokenizer_path = snapshot_dir / "tokenizer.json" + model_file = snapshot_dir / "model.safetensors" + index_path = snapshot_dir / "model.safetensors.index.json" + for p in (config_path, tokenizer_path, model_file, index_path): + if not p.exists(): + raise FileNotFoundError(p) + + config = json.load(open(config_path)) + text_cfg = dict(config["text_config"]) + audio_cfg = dict(config["audio_encoder_config"]) + prompts = args.prompt or list(DEFAULT_PROMPTS) + reference = load_sglang_omni_reference( + args.sglang_omni_src, + required=args.require_sglang_omni_source, + ) + adapter_cls = reference.get("tokenizer_adapter_cls", HiggsTokenizerAdapter) + adapter = load_tokenizer(snapshot_dir, adapter_cls) + prompt_ids = [adapter.build_prompt(p) for p in prompts] + max_len = max(len(x) for x in prompt_ids) + pad_id = int(text_cfg.get("eos_token_id") or 151643) + input_ids = torch.full((len(prompt_ids), max_len), pad_id, dtype=torch.long, device=args.device) + attention_mask = torch.zeros((len(prompt_ids), max_len), dtype=torch.long, device=args.device) + for i, ids in enumerate(prompt_ids): + input_ids[i, : len(ids)] = torch.tensor(ids, dtype=torch.long, device=args.device) + attention_mask[i, : len(ids)] = 1 + prompt_lens = torch.tensor([len(x) for x in prompt_ids], dtype=torch.int64) + + backbone = load_backbone(model_file, text_cfg, args.device) + modality_weight = load_modality_head_weight(model_file, args.device) + trace_tensors: dict[str, torch.Tensor] = {} + trace_hooks = register_trace_hooks(backbone, trace_tensors) if args.trace_out else [] + with torch.inference_mode(): + out = backbone( + input_ids=input_ids, + attention_mask=attention_mask, + use_cache=False, + return_dict=True, + output_hidden_states=bool(args.trace_out), + ) + hidden = out.last_hidden_state + row_idx = torch.arange(len(prompts), device=args.device) + last_hidden = hidden[row_idx, prompt_lens.to(args.device) - 1, :].contiguous() + logits = compute_modality_logits(last_hidden, modality_weight, audio_cfg, reference) + logprobs = torch.log_softmax(logits.to(torch.float32), dim=-1) + top_vals, top_ids = torch.topk(logprobs, k=64, dim=-1) + argmax_ids = torch.argmax(logits, dim=-1).to(torch.int64) + for handle in trace_hooks: + handle.remove() + + tensors = { + "prompt.input_ids_padded": input_ids.cpu().to(torch.int64), + "prompt.attention_mask": attention_mask.cpu().to(torch.int64), + "prompt.lengths": prompt_lens.cpu(), + "final_hidden.bf16": last_hidden.cpu().to(torch.bfloat16), + "audio_logits.f32": logits.cpu().to(torch.float32), + "audio_top64.ids": top_ids.cpu().to(torch.int64), + "audio_top64.logprobs.f32": top_vals.cpu().to(torch.float32), + "audio_argmax.ids": argmax_ids.cpu().to(torch.int64), + } + metadata = { + "fixture_kind": "higgs-one-step-audio-logits-golden", + "schema_version": "1", + "model_id": args.model_id, + "model_revision": args.revision, + "reference": "SGLang-Omni Higgs prompt builder plus Transformers Qwen3 backbone plus SGLang fused modality head semantics", + "sglang_omni_reference_files": "sglang_omni/models/higgs_tts/text_tokenizer.py;sglang_omni/models/higgs_tts/modeling.py;sglang_omni/models/higgs_tts/model.py", + "sglang_omni_source_dir": reference.get("source_dir", ""), + "sglang_omni_source_commit": reference.get("source_commit", ""), + "sglang_omni_direct_imports": "text_tokenizer.py;modeling.py" if reference else "", + "sglang_omni_full_model_imported": "false", + "prompt_count": str(len(prompts)), + "prompts_json": json.dumps(prompts, ensure_ascii=False), + "num_codebooks": str(audio_cfg["num_codebooks"]), + "codebook_vocab_size": str(audio_cfg["vocab_size"]), + "hidden_size": str(text_cfg["hidden_size"]), + "config_sha256": sha256_file(config_path), + "tokenizer_json_sha256": sha256_file(tokenizer_path), + "model_index_sha256": sha256_file(index_path), + "model_safetensors_size": str(model_file.stat().st_size), + "python": platform.python_version(), + "torch": torch.__version__, + "transformers": __import__("transformers").__version__, + "device": torch.cuda.get_device_name(0) if args.device.startswith("cuda") else args.device, + "cuda_peak_allocated": str(torch.cuda.max_memory_allocated() if args.device.startswith("cuda") else 0), + "cuda_peak_reserved": str(torch.cuda.max_memory_reserved() if args.device.startswith("cuda") else 0), + } + out_path = Path(args.out) + out_path.parent.mkdir(parents=True, exist_ok=True) + save_file(tensors, str(out_path), metadata=metadata) + print(f"wrote {out_path} size={out_path.stat().st_size}") + if args.trace_out: + trace_path = Path(args.trace_out) + trace_tensor_count = write_trace_file( + trace_path, + base_tensors=tensors, + trace_tensors=trace_tensors, + hidden_states=out.hidden_states, + prompt_lens=prompt_lens, + last_hidden=last_hidden, + logits=logits, + logprobs=logprobs, + metadata=metadata, + ) + print(f"wrote trace {trace_path} size={trace_path.stat().st_size} tensors={trace_tensor_count}") + print("argmax", argmax_ids.cpu().tolist()) + + +if __name__ == "__main__": + main() diff --git a/tools/higgs/check_higgs_gate_summary.py b/tools/higgs/check_higgs_gate_summary.py new file mode 100755 index 000000000..85370f975 --- /dev/null +++ b/tools/higgs/check_higgs_gate_summary.py @@ -0,0 +1,149 @@ +#!/usr/bin/env python3 +"""Validate a Higgs Audio one-step CUDA gate summary.""" + +from __future__ import annotations + +import argparse +from pathlib import Path + + +REQUIRED_KEYS = { + "status", + "repo", + "commit", + "label", + "model_dir", + "golden", + "sm", + "nvcc_jobs", + "actual", + "session_actual", + "compare_log", + "session_smoke_log", + "session_compare_log", + "auto_view", + "semantic_comparison", + "session_semantic_comparison", + "duplicate_request_id_guard", + "artifacts_nonempty", +} + +OK_KEYS = { + "status", + "semantic_comparison", + "session_semantic_comparison", + "duplicate_request_id_guard", + "artifacts_nonempty", +} + +ARTIFACT_KEYS = ( + "golden", + "actual", + "session_actual", + "compare_log", + "session_smoke_log", + "session_compare_log", +) + +AUTO_VIEW_FILES = ( + "config.json", + "generation_config.json", + "higgs-qwen3-tensor-aliases.json", +) + + +def parse_summary(path: Path) -> dict[str, str]: + values: dict[str, str] = {} + for line_no, raw_line in enumerate(path.read_text().splitlines(), start=1): + line = raw_line.strip() + if not line: + continue + if "=" not in line: + raise ValueError(f"{path}:{line_no}: expected key=value, got {raw_line!r}") + key, value = line.split("=", 1) + if not key: + raise ValueError(f"{path}:{line_no}: empty key") + if key in values: + raise ValueError(f"{path}:{line_no}: duplicate key {key!r}") + values[key] = value + return values + + +def require_equal(values: dict[str, str], key: str, expected: str | None) -> None: + if expected is not None and values.get(key) != expected: + raise ValueError(f"{key}={values.get(key)!r}, expected {expected!r}") + + +def require_nonempty_file(path: Path) -> None: + if not path.is_file() or path.stat().st_size == 0: + raise ValueError(f"required artifact missing or empty: {path}") + + +def validate(args: argparse.Namespace) -> dict[str, str]: + summary = args.summary + if not summary.is_file() or summary.stat().st_size == 0: + raise ValueError(f"summary missing or empty: {summary}") + + values = parse_summary(summary) + missing = sorted(REQUIRED_KEYS - values.keys()) + if missing: + raise ValueError(f"summary missing required keys: {', '.join(missing)}") + + extra = sorted(values.keys() - REQUIRED_KEYS) + if extra and not args.allow_extra: + raise ValueError(f"summary has unknown keys: {', '.join(extra)}") + + for key in OK_KEYS: + if values[key] != "ok": + raise ValueError(f"{key}={values[key]!r}, expected 'ok'") + + if not values["commit"]: + raise ValueError("commit is empty") + + if not values["sm"].isdigit(): + raise ValueError(f"sm must be numeric, got {values['sm']!r}") + + if not values["nvcc_jobs"].isdigit(): + raise ValueError(f"nvcc_jobs must be numeric, got {values['nvcc_jobs']!r}") + + require_equal(values, "label", args.expected_label) + require_equal(values, "sm", args.expected_sm) + require_equal(values, "nvcc_jobs", args.expected_nvcc_jobs) + require_equal(values, "model_dir", args.expected_model_dir) + require_equal(values, "golden", args.expected_golden) + + if args.check_files: + for key in ARTIFACT_KEYS: + require_nonempty_file(Path(values[key])) + auto_view = Path(values["auto_view"]) + if not auto_view.is_dir(): + raise ValueError(f"auto_view is not a directory: {auto_view}") + for name in AUTO_VIEW_FILES: + require_nonempty_file(auto_view / name) + + return values + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("summary", type=Path) + parser.add_argument("--expected-label") + parser.add_argument("--expected-sm") + parser.add_argument("--expected-nvcc-jobs") + parser.add_argument("--expected-model-dir") + parser.add_argument("--expected-golden") + parser.add_argument("--check-files", action="store_true") + parser.add_argument("--allow-extra", action="store_true") + args = parser.parse_args() + + values = validate(args) + print( + "higgs gate summary: ok " + f"commit={values['commit']} " + f"label={values['label']} " + f"sm={values['sm']}" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/higgs/check_higgs_sglang_omni_imports.py b/tools/higgs/check_higgs_sglang_omni_imports.py new file mode 100755 index 000000000..3986875ba --- /dev/null +++ b/tools/higgs/check_higgs_sglang_omni_imports.py @@ -0,0 +1,182 @@ +#!/usr/bin/env python3 +"""Probe the SGLang-Omni Higgs import boundary. + +This is intentionally narrower than a serving/runtime parity gate. The direct +Higgs modules are the source contract used by the one-step golden generator; +the full model import checks whether the local environment can even start the +SGLang-Omni runtime path. +""" + +from __future__ import annotations + +import argparse +import importlib +import importlib.metadata +import subprocess +import sys +import tomllib +from pathlib import Path + + +DIRECT_MODULES = ( + "sglang_omni.models.higgs_tts.text_tokenizer", + "sglang_omni.models.higgs_tts.modeling", + "sglang_omni.models.higgs_tts.hf_config", +) +FULL_MODEL_MODULE = "sglang_omni.models.higgs_tts.model" +TORCH_POOL_API = "_cuda_beginAllocateCurrentThreadToPool" + + +def git_short_commit(path: Path) -> str: + result = subprocess.run( + ["git", "-C", str(path), "rev-parse", "--short", "HEAD"], + check=False, + capture_output=True, + text=True, + ) + return result.stdout.strip() if result.returncode == 0 else "unknown" + + +def module_status(name: str) -> tuple[str, str]: + try: + importlib.import_module(name) + except Exception as exc: + return "fail", f"{type(exc).__name__}:{exc}" + return "ok", "" + + +def package_version(name: str) -> str: + try: + return importlib.metadata.version(name) + except importlib.metadata.PackageNotFoundError: + return "missing" + + +def torch_stack() -> dict[str, str]: + try: + import torch + import torch.cuda.memory as torch_cuda_memory + except Exception as exc: + reason = normalize_reason(f"{type(exc).__name__}:{exc}") + return { + "package.torch.version": package_version("torch"), + "package.torch.cuda": "unknown", + "package.torch.has_cuda_begin_allocate_current_thread_to_pool": "fail", + "package.torch.import_reason": reason, + } + + return { + "package.torch.version": str(torch.__version__), + "package.torch.cuda": str(torch.version.cuda), + "package.torch.has_cuda_begin_allocate_current_thread_to_pool": "ok" + if hasattr(torch_cuda_memory, TORCH_POOL_API) + else "fail", + } + + +def normalize_dependency_name(spec: str) -> str: + name = spec.split(";", 1)[0].strip() + for sep in ("[", "<", ">", "=", "!", "~"): + name = name.split(sep, 1)[0].strip() + return name.replace("_", "-").lower() + + +def pyproject_requirements(src: Path) -> dict[str, str]: + pyproject = src / "pyproject.toml" + if not pyproject.is_file(): + return {} + data = tomllib.loads(pyproject.read_text()) + project = data.get("project", {}) + deps = project.get("dependencies", []) + by_name = {normalize_dependency_name(dep): dep for dep in deps} + wanted = ( + "torch", + "sglang", + "transformers", + "flash-attn-4", + "flashinfer-python", + "nvidia-cutlass-dsl", + ) + out = { + "pyproject.requires_python": str(project.get("requires-python", "unknown")), + } + for name in wanted: + out[f"pyproject.dependency.{name}"] = by_name.get(name, "missing") + return out + + +def normalize_reason(reason: str) -> str: + if "No module named 'sglang'" in reason or 'No module named "sglang"' in reason: + return "missing_sglang" + return reason.replace("\n", " ") + + +def print_kv(key: str, value: str) -> None: + print(f"{key}={value}") + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument( + "--sglang-omni-src", + required=True, + type=Path, + help="SGLang-Omni source tree containing sglang_omni/.", + ) + parser.add_argument( + "--require-direct", + action="store_true", + help="Return nonzero if direct Higgs source modules cannot be imported.", + ) + parser.add_argument( + "--require-full-model", + action="store_true", + help="Return nonzero if the full SGLang-Omni Higgs model cannot be imported.", + ) + args = parser.parse_args() + + src = args.sglang_omni_src.resolve() + if not (src / "sglang_omni/models/higgs_tts").is_dir(): + raise SystemExit(f"SGLang-Omni Higgs source not found: {src}") + + sys.path.insert(0, str(src)) + + print_kv("sglang_omni_src", str(src)) + print_kv("sglang_omni_commit", git_short_commit(src)) + print_kv("python.executable", sys.executable) + print_kv("python.version", sys.version.replace("\n", " ")) + for key, value in pyproject_requirements(src).items(): + print_kv(key, value) + for key, value in torch_stack().items(): + print_kv(key, value) + print_kv("package.sglang.version", package_version("sglang")) + print_kv("package.transformers.version", package_version("transformers")) + print_kv("package.sgl-kernel.version", package_version("sgl-kernel")) + + direct_ok = True + for name in DIRECT_MODULES: + status, reason = module_status(name) + direct_ok = direct_ok and status == "ok" + print_kv(f"module.{name}", status) + if reason: + print_kv(f"module.{name}.reason", normalize_reason(reason)) + + full_status, full_reason = module_status(FULL_MODEL_MODULE) + print_kv(f"module.{FULL_MODEL_MODULE}", full_status) + if full_reason: + print_kv(f"module.{FULL_MODEL_MODULE}.reason", normalize_reason(full_reason)) + + print_kv("direct_higgs_imports", "ok" if direct_ok else "fail") + if full_status == "ok": + print_kv("full_higgs_model_import", "ok") + else: + print_kv("full_higgs_model_import", normalize_reason(full_reason)) + + if args.require_direct and not direct_ok: + raise SystemExit(1) + if args.require_full_model and full_status != "ok": + raise SystemExit(1) + + +if __name__ == "__main__": + main() diff --git a/tools/higgs/check_higgs_sglang_omni_runtime_readiness_summary.py b/tools/higgs/check_higgs_sglang_omni_runtime_readiness_summary.py new file mode 100755 index 000000000..23b2ed272 --- /dev/null +++ b/tools/higgs/check_higgs_sglang_omni_runtime_readiness_summary.py @@ -0,0 +1,195 @@ +#!/usr/bin/env python3 +"""Validate a Higgs Audio SGLang-Omni full-runtime readiness summary.""" + +from __future__ import annotations + +import argparse +from pathlib import Path + + +REQUIRED_KEYS = { + "status", + "repo", + "commit", + "label", + "python", + "python_version", + "sglang_omni_src", + "sglang_omni_commit", + "readiness_log", + "pyproject_torch", + "pyproject_sglang", + "torch_version", + "torch_cuda", + "torch_has_cuda_pool_api", + "sglang_version", + "transformers_version", + "sglang_omni_direct_imports", + "sglang_omni_full_model_import", + "runtime_ready", + "artifacts_nonempty", +} + +OK_KEYS = { + "status", + "sglang_omni_direct_imports", + "sglang_omni_full_model_import", + "runtime_ready", + "artifacts_nonempty", +} + + +def parse_summary(path: Path) -> dict[str, str]: + values: dict[str, str] = {} + for line_no, raw_line in enumerate(path.read_text().splitlines(), start=1): + line = raw_line.strip() + if not line: + continue + if "=" not in line: + raise ValueError(f"{path}:{line_no}: expected key=value, got {raw_line!r}") + key, value = line.split("=", 1) + if not key: + raise ValueError(f"{path}:{line_no}: empty key") + if key in values: + raise ValueError(f"{path}:{line_no}: duplicate key {key!r}") + values[key] = value + return values + + +def require_equal(values: dict[str, str], key: str, expected: str | None) -> None: + if expected is not None and values.get(key) != expected: + raise ValueError(f"{key}={values.get(key)!r}, expected {expected!r}") + + +def require_nonempty_file(path: Path) -> None: + if not path.is_file() or path.stat().st_size == 0: + raise ValueError(f"required artifact missing or empty: {path}") + + +def require_log_line(path: Path, expected: str) -> None: + lines = {line.strip() for line in path.read_text().splitlines()} + if expected not in lines: + raise ValueError(f"{path} missing expected line: {expected}") + + +def validate(args: argparse.Namespace) -> dict[str, str]: + summary = args.summary + if not summary.is_file() or summary.stat().st_size == 0: + raise ValueError(f"summary missing or empty: {summary}") + + values = parse_summary(summary) + missing = sorted(REQUIRED_KEYS - values.keys()) + if missing: + raise ValueError(f"summary missing required keys: {', '.join(missing)}") + + extra = sorted(values.keys() - REQUIRED_KEYS) + if extra and not args.allow_extra: + raise ValueError(f"summary has unknown keys: {', '.join(extra)}") + + require_equal(values, "status", args.expected_status) + require_equal(values, "label", args.expected_label) + require_equal(values, "python", args.expected_python) + require_equal(values, "sglang_omni_src", args.expected_sglang_omni_src) + require_equal(values, "sglang_omni_commit", args.expected_sglang_omni_commit) + + if values["status"] not in {"ok", "fail"}: + raise ValueError(f"status must be 'ok' or 'fail', got {values['status']!r}") + + if values["runtime_ready"] not in {"ok", "fail"}: + raise ValueError( + f"runtime_ready must be 'ok' or 'fail', got {values['runtime_ready']!r}" + ) + + if values["status"] == "ok": + for key in OK_KEYS: + if values[key] != "ok": + raise ValueError(f"{key}={values[key]!r}, expected 'ok'") + + if not values["commit"]: + raise ValueError("commit is empty") + + if not values["sglang_omni_commit"]: + raise ValueError("sglang_omni_commit is empty") + + nonempty_keys = ( + "python_version", + "pyproject_torch", + "pyproject_sglang", + "torch_version", + "torch_cuda", + "torch_has_cuda_pool_api", + "sglang_version", + "transformers_version", + ) + for key in nonempty_keys: + if not values[key]: + raise ValueError(f"{key} is empty") + + if values["status"] == "ok" and values["sglang_omni_direct_imports"] != "ok": + raise ValueError("status=ok requires sglang_omni_direct_imports=ok") + + if values["runtime_ready"] == "ok" and values["sglang_omni_direct_imports"] != "ok": + raise ValueError("runtime_ready=ok requires sglang_omni_direct_imports=ok") + + if values["runtime_ready"] == "ok" and values["sglang_omni_full_model_import"] != "ok": + raise ValueError("runtime_ready=ok requires sglang_omni_full_model_import=ok") + + if values["status"] == "ok" and values["runtime_ready"] != "ok": + raise ValueError("status=ok requires runtime_ready=ok") + + if values["status"] == "fail" and values["runtime_ready"] == "ok": + raise ValueError("status=fail cannot report runtime_ready=ok") + + if args.check_files: + readiness_log = Path(values["readiness_log"]) + require_nonempty_file(readiness_log) + log_mirrors = { + "python.version": values["python_version"], + "pyproject.dependency.torch": values["pyproject_torch"], + "pyproject.dependency.sglang": values["pyproject_sglang"], + "package.torch.version": values["torch_version"], + "package.torch.cuda": values["torch_cuda"], + "package.torch.has_cuda_begin_allocate_current_thread_to_pool": values[ + "torch_has_cuda_pool_api" + ], + "package.sglang.version": values["sglang_version"], + "package.transformers.version": values["transformers_version"], + } + for key, value in log_mirrors.items(): + require_log_line(readiness_log, f"{key}={value}") + require_log_line( + readiness_log, + f"direct_higgs_imports={values['sglang_omni_direct_imports']}", + ) + require_log_line( + readiness_log, + f"full_higgs_model_import={values['sglang_omni_full_model_import']}", + ) + + return values + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("summary", type=Path) + parser.add_argument("--expected-status") + parser.add_argument("--expected-label") + parser.add_argument("--expected-python") + parser.add_argument("--expected-sglang-omni-src") + parser.add_argument("--expected-sglang-omni-commit") + parser.add_argument("--check-files", action="store_true") + parser.add_argument("--allow-extra", action="store_true") + args = parser.parse_args() + + values = validate(args) + print( + "higgs sglang-omni runtime readiness summary: ok " + f"status={values['status']} " + f"label={values['label']} " + f"sglang_omni_commit={values['sglang_omni_commit']} " + f"full_model_import={values['sglang_omni_full_model_import']}" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/higgs/check_higgs_sglang_omni_source_gate_summary.py b/tools/higgs/check_higgs_sglang_omni_source_gate_summary.py new file mode 100755 index 000000000..601e32158 --- /dev/null +++ b/tools/higgs/check_higgs_sglang_omni_source_gate_summary.py @@ -0,0 +1,144 @@ +#!/usr/bin/env python3 +"""Validate a Higgs Audio SGLang-Omni source-reference gate summary.""" + +from __future__ import annotations + +import argparse +from pathlib import Path + + +REQUIRED_KEYS = { + "status", + "repo", + "commit", + "label", + "model_dir", + "sglang_omni_src", + "sglang_omni_commit", + "golden", + "reference", + "compare_log", + "readiness_log", + "sglang_omni_direct_imports", + "sglang_omni_full_model_import", + "source_reference_strict_comparison", + "artifacts_nonempty", +} + +OK_KEYS = { + "status", + "sglang_omni_direct_imports", + "source_reference_strict_comparison", + "artifacts_nonempty", +} + +ARTIFACT_KEYS = ( + "golden", + "reference", + "compare_log", + "readiness_log", +) + +def parse_summary(path: Path) -> dict[str, str]: + values: dict[str, str] = {} + for line_no, raw_line in enumerate(path.read_text().splitlines(), start=1): + line = raw_line.strip() + if not line: + continue + if "=" not in line: + raise ValueError(f"{path}:{line_no}: expected key=value, got {raw_line!r}") + key, value = line.split("=", 1) + if not key: + raise ValueError(f"{path}:{line_no}: empty key") + if key in values: + raise ValueError(f"{path}:{line_no}: duplicate key {key!r}") + values[key] = value + return values + + +def require_equal(values: dict[str, str], key: str, expected: str | None) -> None: + if expected is not None and values.get(key) != expected: + raise ValueError(f"{key}={values.get(key)!r}, expected {expected!r}") + + +def require_nonempty_file(path: Path) -> None: + if not path.is_file() or path.stat().st_size == 0: + raise ValueError(f"required artifact missing or empty: {path}") + + +def require_log_line(path: Path, expected: str) -> None: + lines = {line.strip() for line in path.read_text().splitlines()} + if expected not in lines: + raise ValueError(f"{path} missing expected line: {expected}") + + +def validate(args: argparse.Namespace) -> dict[str, str]: + summary = args.summary + if not summary.is_file() or summary.stat().st_size == 0: + raise ValueError(f"summary missing or empty: {summary}") + + values = parse_summary(summary) + missing = sorted(REQUIRED_KEYS - values.keys()) + if missing: + raise ValueError(f"summary missing required keys: {', '.join(missing)}") + + extra = sorted(values.keys() - REQUIRED_KEYS) + if extra and not args.allow_extra: + raise ValueError(f"summary has unknown keys: {', '.join(extra)}") + + for key in OK_KEYS: + if values[key] != "ok": + raise ValueError(f"{key}={values[key]!r}, expected 'ok'") + + if not values["commit"]: + raise ValueError("commit is empty") + + if not values["sglang_omni_commit"]: + raise ValueError("sglang_omni_commit is empty") + + if not values["sglang_omni_full_model_import"]: + raise ValueError("sglang_omni_full_model_import is empty") + + require_equal(values, "label", args.expected_label) + require_equal(values, "model_dir", args.expected_model_dir) + require_equal(values, "sglang_omni_src", args.expected_sglang_omni_src) + require_equal(values, "sglang_omni_commit", args.expected_sglang_omni_commit) + require_equal(values, "golden", args.expected_golden) + + if args.check_files: + for key in ARTIFACT_KEYS: + require_nonempty_file(Path(values[key])) + require_log_line(Path(values["compare_log"]), "higgs one-step strict comparison: ok") + require_log_line(Path(values["readiness_log"]), "direct_higgs_imports=ok") + require_log_line( + Path(values["readiness_log"]), + f"full_higgs_model_import={values['sglang_omni_full_model_import']}", + ) + + return values + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("summary", type=Path) + parser.add_argument("--expected-label") + parser.add_argument("--expected-model-dir") + parser.add_argument("--expected-sglang-omni-src") + parser.add_argument("--expected-sglang-omni-commit") + parser.add_argument("--expected-golden") + parser.add_argument("--check-files", action="store_true") + parser.add_argument("--allow-extra", action="store_true") + args = parser.parse_args() + + values = validate(args) + print( + "higgs sglang-omni source gate summary: ok " + f"commit={values['commit']} " + f"label={values['label']} " + f"sglang_omni_commit={values['sglang_omni_commit']} " + f"full_model_import={values['sglang_omni_full_model_import']}" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/higgs/run_higgs_one_step_cuda_gate.sh b/tools/higgs/run_higgs_one_step_cuda_gate.sh new file mode 100755 index 000000000..46bf543b4 --- /dev/null +++ b/tools/higgs/run_higgs_one_step_cuda_gate.sh @@ -0,0 +1,250 @@ +#!/usr/bin/env bash +set -euo pipefail + +usage() { + cat <<'USAGE' +Run the Higgs Audio one-step CUDA golden gate. + +Required: + --model-dir DIR Higgs checkpoint directory. + +Optional: + --golden FILE Golden safetensors fixture. + Default: test_data/higgs-one-step-audio-logits.safetensors + --result-root DIR Result directory. + Default: /data/results/pegainfer/higgs-audio + --label LABEL Output label. Default: current git short SHA. + --sm SM PEGAINFER_CUDA_SM. Default: 89 + --nvcc-jobs N PEGAINFER_NVCC_JOBS. Default: 8 + --profile Also capture an NSYS profile for the actual dump. + -h, --help Show this help. + +Outputs: + /actual/higgs-one-step-actual-cuda-bf16-auto-