From cf5c35b6fff0fad0fb280de917665d6f0a057e77 Mon Sep 17 00:00:00 2001 From: Adam Page Date: Tue, 21 Jul 2026 18:26:11 -0500 Subject: [PATCH 1/4] feat: publish multi-precision RVQ encoders --- .devcontainer/devcontainer.json | 8 +- .devcontainer/install.sh | 21 +- compressionkit/configs/ecg_rvq.py | 16 +- compressionkit/configs/ppg_rvq.py | 16 +- compressionkit/experiments/repackage.py | 126 ++++++++- compressionkit/export/__init__.py | 37 +++ compressionkit/export/artifact_contract.py | 27 ++ compressionkit/export/demo_ecg.py | 299 +++++++++++++++++++++ compressionkit/export/demo_ppg.py | 205 ++++++++++++++ compressionkit/export/demo_recordings.py | 136 ++++++++++ compressionkit/export/deploy.py | 107 +++++++- compressionkit/export/family_registry.py | 9 +- compressionkit/export/model_card.py | 40 +++ compressionkit/export/quantization.py | 242 +++++++++++++++++ compressionkit/export/validate.py | 61 +++++ compressionkit/preprocessing/ecg.py | 12 +- compressionkit/preprocessing/ppg.py | 12 +- compressionkit/recipes/base_rvq.py | 11 +- compressionkit/trainers/common.py | 63 +++++ configs/ppg_rvq_smoketest.yaml | 3 + docs/deployment.md | 31 +++ docs/huggingface.md | 23 +- docs/methods/rvq.md | 1 + docs/release-contract.md | 1 + scripts/attach_rvq_demo_recordings.py | 141 ++++++++++ scripts/compare_rvq_encoder_precisions.py | 166 ++++++++++++ scripts/devcontainer.sh | 72 ++++- scripts/generate_ecg_demo_samples.py | 51 ++++ scripts/publish_to_huggingface.py | 2 + tests/test_base_rvq_scorecard_sync.py | 12 +- tests/test_demo_ecg.py | 72 +++++ tests/test_demo_ppg.py | 39 +++ tests/test_demo_recordings.py | 57 ++++ tests/test_deploy_enhanced.py | 6 + tests/test_huggingface.py | 32 +++ tests/test_rvq_deploy_smoke.py | 27 ++ tests/test_rvq_quantization.py | 95 +++++++ 37 files changed, 2239 insertions(+), 40 deletions(-) create mode 100644 compressionkit/export/demo_ecg.py create mode 100644 compressionkit/export/demo_ppg.py create mode 100644 compressionkit/export/demo_recordings.py create mode 100644 compressionkit/export/quantization.py create mode 100644 scripts/attach_rvq_demo_recordings.py create mode 100644 scripts/compare_rvq_encoder_precisions.py create mode 100644 scripts/generate_ecg_demo_samples.py create mode 100644 tests/test_demo_ecg.py create mode 100644 tests/test_demo_ppg.py create mode 100644 tests/test_demo_recordings.py create mode 100644 tests/test_rvq_quantization.py diff --git a/.devcontainer/devcontainer.json b/.devcontainer/devcontainer.json index ab7528d..5d91b94 100644 --- a/.devcontainer/devcontainer.json +++ b/.devcontainer/devcontainer.json @@ -36,18 +36,18 @@ // ], "mounts": [ - "source=/data/datasets,target=/home/vscode/datasets,type=bind,consistency=cached" + "source=/data/datasets,target=/home/vscode/datasets,type=bind,consistency=cached", + "source=compressionkit-huggingface,target=/home/vscode/.cache/huggingface,type=volume" ], "forwardPorts": [6006], - "postCreateCommand": "./.devcontainer/install.sh", + "postCreateCommand": "sudo install -d -o vscode -g vscode -m 700 /home/vscode/.cache /home/vscode/.cache/huggingface && ./.devcontainer/install.sh", "remoteEnv": { "LD_LIBRARY_PATH": "${containerEnv:LD_LIBRARY_PATH}:/usr/local/cuda/lib64", "PATH": "${containerEnv:PATH}:/usr/local/cuda/bin", - "TF_FORCE_GPU_ALLOW_GROWTH": "true", - "HF_TOKEN": "${localEnv:HF_TOKEN}" + "TF_FORCE_GPU_ALLOW_GROWTH": "true" }, "customizations": { diff --git a/.devcontainer/install.sh b/.devcontainer/install.sh index 9a91afe..b44462b 100755 --- a/.devcontainer/install.sh +++ b/.devcontainer/install.sh @@ -1,9 +1,24 @@ #!/bin/bash +set -euo pipefail + export DEBIAN_FRONTEND=noninteractive -sudo apt update -sudo apt install -y \ +apt_retry() { + local attempts=0 + until sudo apt-get "$@"; do + attempts=$((attempts + 1)) + if (( attempts >= 24 )); then + echo "apt-get failed after ${attempts} attempts: $*" >&2 + return 1 + fi + echo "apt-get is unavailable; retrying in 5 seconds (${attempts}/24)..." >&2 + sleep 5 + done +} + +apt_retry update +apt_retry install -y \ build-essential \ cmake \ make \ @@ -15,7 +30,7 @@ sudo apt install -y \ git-lfs # Optional: reduce image size a bit -sudo apt clean +sudo apt-get clean sudo rm -rf /var/lib/apt/lists/* git lfs install diff --git a/compressionkit/configs/ecg_rvq.py b/compressionkit/configs/ecg_rvq.py index 3871d61..c6c3698 100644 --- a/compressionkit/configs/ecg_rvq.py +++ b/compressionkit/configs/ecg_rvq.py @@ -2,7 +2,7 @@ from __future__ import annotations -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator from compressionkit.configs.paths import default_datasets_dir @@ -483,10 +483,24 @@ class EvaluationConfig(BaseModel): Capped at ``num_samples``. Plots are slow to render and large to ship, so this is typically much smaller than ``num_samples``.""" tflite_rep_batches: int = 8 + """Legacy representative-batch setting; retained for config compatibility.""" + int8_calibration_frames: int = Field(default=4096, gt=0) + """Real normalized validation frames used to calibrate INT8 LiteRT ranges.""" + int8_validation_frames: int = Field(default=2048, gt=0) + """Disjoint real validation frames used to measure INT8 encoder parity.""" + int8_sampling_pool_frames: int = Field(default=65_536, gt=0) + """Maximum validation frames reservoir-sampled for quantization partitions.""" input_bit_depth: int = 16 band_metrics: BandMetricsConfig = Field(default_factory=BandMetricsConfig) stitching: StitchingEvalConfig = Field(default_factory=StitchingEvalConfig) + @model_validator(mode="after") + def validate_quantization_partition(self) -> EvaluationConfig: + """Require the sampling pool to cover both disjoint partitions.""" + if self.int8_sampling_pool_frames < self.int8_calibration_frames + self.int8_validation_frames: + raise ValueError("int8_sampling_pool_frames must cover calibration plus validation frames") + return self + class WandbConfig(BaseModel): """Weights & Biases logging configuration.""" diff --git a/compressionkit/configs/ppg_rvq.py b/compressionkit/configs/ppg_rvq.py index 3ce8ab6..77903fd 100644 --- a/compressionkit/configs/ppg_rvq.py +++ b/compressionkit/configs/ppg_rvq.py @@ -2,7 +2,7 @@ from __future__ import annotations -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator from compressionkit.configs.artifact_suite import ArtifactSuiteConfig from compressionkit.configs.paths import default_datasets_dir @@ -311,11 +311,25 @@ class EvaluationConfig(BaseModel): num_plot_samples: int = 50 """Subset of ``num_samples`` that also receive a PNG plot artifact.""" tflite_rep_batches: int = 8 + """Legacy representative-batch setting; retained for config compatibility.""" + int8_calibration_frames: int = Field(default=4096, gt=0) + """Real normalized validation frames used to calibrate INT8 LiteRT ranges.""" + int8_validation_frames: int = Field(default=2048, gt=0) + """Disjoint real validation frames used to measure INT8 encoder parity.""" + int8_sampling_pool_frames: int = Field(default=65_536, gt=0) + """Maximum validation frames reservoir-sampled for quantization partitions.""" input_bit_depth: int = 16 band_metrics: BandMetricsConfig = Field(default_factory=BandMetricsConfig) physiokit_metrics: PhysiokitMetricsConfig = Field(default_factory=PhysiokitMetricsConfig) long_recording: LongRecordingEvalConfig = Field(default_factory=LongRecordingEvalConfig) + @model_validator(mode="after") + def validate_quantization_partition(self) -> EvaluationConfig: + """Require the sampling pool to cover both disjoint partitions.""" + if self.int8_sampling_pool_frames < self.int8_calibration_frames + self.int8_validation_frames: + raise ValueError("int8_sampling_pool_frames must cover calibration plus validation frames") + return self + class WandbConfig(BaseModel): """Weights & Biases logging configuration.""" diff --git a/compressionkit/experiments/repackage.py b/compressionkit/experiments/repackage.py index 3a1bc05..08d2b29 100644 --- a/compressionkit/experiments/repackage.py +++ b/compressionkit/experiments/repackage.py @@ -41,6 +41,88 @@ def finalize_release_metadata(output_dir: Path, scorecard_payload: dict[str, obj write_checksums(output_dir) +def collect_release_quantization_frames( + experiment: GoldenExperiment, + run_dir: Path, +) -> tuple[np.ndarray, np.ndarray, dict[str, object]]: + """Rebuild disjoint real, training-preprocessed INT8 frame partitions. + + The release repackage path must not calibrate a model trained on normalized + physiological windows with raw synthetic waveforms. This helper reloads the + golden's saved Pydantic config and delegates to its normal dataset builder, + so resampling, framing, and normalization remain identical to training. + """ + from compressionkit.trainers.common import collect_disjoint_quantization_datasets + + config_path = run_dir / "config.json" + if not config_path.is_file(): + raise FileNotFoundError(f"Missing saved golden config: {config_path}") + config_payload = config_path.read_text() + + if experiment.modality == "ppg": + from compressionkit.configs.ppg_rvq import PpgRvqConfig + from compressionkit.preprocessing.ppg import build_augmenter, build_preprocessor + from compressionkit.trainers.ppg_rvq import build_datasets + + cfg = PpgRvqConfig.model_validate_json(config_payload) + preprocessor = build_preprocessor( + frame_size=cfg.data.frame_size, + epsilon=cfg.data.epsilon, + seed=cfg.data.shuffle_seed, + ) + augmenter = build_augmenter(tuple(cfg.data.gaussian_noise), aug_cfg=cfg.data.augmentation) + _, val_ds, _, _ = build_datasets(cfg, preprocessor, augmenter) + sample_rate = cfg.data.sampling_rate + elif experiment.modality == "ecg": + from compressionkit.configs.ecg_rvq import EcgRvqConfig + from compressionkit.preprocessing.ecg import build_augmenter, build_preprocessor + from compressionkit.trainers.ecg_rvq import build_datasets + + cfg = EcgRvqConfig.model_validate_json(config_payload) + preprocessor = build_preprocessor( + frame_size=cfg.data.frame_size, + epsilon=cfg.data.epsilon, + seed=cfg.data.shuffle_seed, + ) + augmenter = build_augmenter( + aug_cfg=cfg.data.augmentation, + sample_rate=cfg.data.effective_sample_rate, + seed=cfg.data.shuffle_seed, + ) + _, val_ds, _, _ = build_datasets(cfg, preprocessor, augmenter) + sample_rate = cfg.data.effective_sample_rate + else: + raise ValueError(f"Unsupported RVQ modality {experiment.modality!r}") + + calibration_frames, validation_frames = collect_disjoint_quantization_datasets( + val_ds, + calibration_frames=cfg.evaluation.int8_calibration_frames, + validation_frames=cfg.evaluation.int8_validation_frames, + sampling_pool_frames=cfg.evaluation.int8_sampling_pool_frames, + seed=cfg.data.shuffle_seed, + ) + contract: dict[str, object] = { + "format_version": 1, + "sample_rate_hz": sample_rate, + "frame_size": cfg.data.frame_size, + "input_shape": [1, 1, cfg.data.frame_size, 1], + "normalization": { + "kind": "per_frame_layer_norm", + "mean": "mean over all samples in each frame", + "variance": "mean squared deviation over all samples in each frame", + "epsilon": cfg.data.epsilon, + "inverse_for_display": "raw = normalized * sqrt(variance + epsilon) + mean", + }, + "int8_quantization": { + "calibration_frames": int(calibration_frames.shape[0]), + "validation_frames": int(validation_frames.shape[0]), + "sampling_pool_frames": cfg.evaluation.int8_sampling_pool_frames, + "sampling_method": "seeded reservoir sample; disjoint calibration and validation partitions", + }, + } + return calibration_frames, validation_frames, contract + + def repackage_rvq_golden( experiment: GoldenExperiment, *, @@ -57,7 +139,7 @@ def repackage_rvq_golden( import keras from compressionkit.export.deploy import export_for_deployment - from compressionkit.export.stimulus import export_stimulus_npz, generate_stimulus + from compressionkit.export.stimulus import export_stimulus_npz output_dir = output_dir or (run_dir / "deploy") resolved_scorecard_path = resolve_scorecard_path(run_dir, scorecard_path) @@ -82,17 +164,13 @@ def repackage_rvq_golden( else: frame_size = input_shape[-1] - rep_dataset = generate_stimulus( - modality=experiment.modality, - num_samples=100, - frame_size=frame_size, - sample_rate=experiment.sample_rate, - seed=42, + rep_dataset, quantization_validation_dataset, preprocessing_contract = collect_release_quantization_frames( + experiment, run_dir ) - if len(input_shape) == 4: - rep_dataset = rep_dataset.reshape(-1, 1, frame_size, 1) - elif len(input_shape) == 3: - rep_dataset = rep_dataset[..., np.newaxis] + if tuple(rep_dataset.shape[1:]) != tuple(input_shape[1:]): + raise ValueError( + f"Calibration frame shape {rep_dataset.shape[1:]} does not match encoder input {input_shape[1:]}" + ) model_card_info: dict[str, object] = { "experiment_id": experiment.experiment_id, @@ -101,7 +179,24 @@ def repackage_rvq_golden( "sample_rate": experiment.sample_rate, "compression_ratio": experiment.compression_ratio, "license": "other", + "preprocessing_contract": preprocessing_contract, } + precision_paths = [ + run_dir.parent / "precision-study" / f"{experiment.experiment_id}.json", + run_dir.parent / "rvq_encoder_precision_comparison.json", + ] + for precision_path in precision_paths: + if not precision_path.is_file(): + continue + precision_payload = json.loads(precision_path.read_text()) + variants = precision_payload.get("experiments", {}).get(experiment.experiment_id, {}).get("variants", {}) + if isinstance(variants, dict): + model_card_info["encoder_precision_report"] = { + name: payload.get("report", {}) + for name, payload in variants.items() + if isinstance(payload, dict) + } + break if scorecard_payload is not None: model_card_info["scorecard_summary"] = scorecard_payload @@ -116,6 +211,7 @@ def repackage_rvq_golden( decoder=decoder, rvq_weights=rvq_weights, rep_dataset=rep_dataset, + quantization_validation_dataset=quantization_validation_dataset, output_dir=output_dir, sample_inputs=sample_inputs, sample_targets=sample_inputs, @@ -125,6 +221,8 @@ def repackage_rvq_golden( export_decoder_int8=export_decoder_int8, model_card_info=model_card_info, ) + quantization_report = json.loads((output_dir / "quantization_report.json").read_text()) + report_metrics = quantization_report["metrics"] export_stimulus_npz( modality=experiment.modality, output_path=output_dir / "sample_stimulus.npz", @@ -140,10 +238,16 @@ def repackage_rvq_golden( "deploy_dir": str(output_dir), "scorecard_path": str(resolved_scorecard_path) if resolved_scorecard_path is not None else None, "artifacts": artifacts.as_dict(), + "quantization_report": { + "passed": quantization_report["passed"], + "reconstruction_prd_percent_p90": report_metrics["reconstruction_prd_percent_p90"], + "encoder_input_saturation_fraction_max": report_metrics["encoder_input_saturation_fraction_max"], + }, } __all__ = [ + "collect_release_quantization_frames", "finalize_release_metadata", "load_scorecard_payload", "repackage_rvq_golden", diff --git a/compressionkit/export/__init__.py b/compressionkit/export/__init__.py index 2ca80b0..2efecda 100644 --- a/compressionkit/export/__init__.py +++ b/compressionkit/export/__init__.py @@ -5,7 +5,27 @@ export_codebooks_npz, extract_codebooks, ) +from compressionkit.export.demo_ecg import ( + EcgDemoClip, + EcgDemoExport, + EcgDemoQuality, + assess_ecg_demo_clip, + export_ecg_demo_csvs, + generate_ecg_demo_clips, +) +from compressionkit.export.demo_ppg import PpgDemoClip, PpgDemoQuality, assess_ppg_demo_clip, generate_ppg_demo_clips +from compressionkit.export.demo_recordings import ( + DemoRecordingsExport, + attach_demo_recordings_to_deploy, + export_demo_recordings, +) from compressionkit.export.deploy import DeploymentArtifacts, export_for_deployment, sync_scorecard_to_deploy +from compressionkit.export.quantization import ( + RvqEncoderQuantizationReport, + evaluate_rvq_encoder_quantization, + require_rvq_encoder_quantization_report, + write_rvq_encoder_quantization_report, +) from compressionkit.export.release import ( build_model_card, build_release_metadata, @@ -24,25 +44,42 @@ from compressionkit.export.validate import DeployValidationResult, validate_deploy_package __all__ = [ + "DemoRecordingsExport", "DeployValidationResult", "DeploymentArtifacts", + "EcgDemoClip", + "EcgDemoExport", + "EcgDemoQuality", + "PpgDemoClip", + "PpgDemoQuality", + "RvqEncoderQuantizationReport", "SpihtDeploymentArtifacts", + "assess_ecg_demo_clip", + "assess_ppg_demo_clip", + "attach_demo_recordings_to_deploy", "build_model_card", "build_release_metadata", + "evaluate_rvq_encoder_quantization", "export_codebooks_header", "export_codebooks_npz", "export_decoder_tflite", + "export_demo_recordings", + "export_ecg_demo_csvs", "export_encoder_tflite", "export_for_deployment", "export_spiht_deploy", "export_stimulus_npz", "extract_codebooks", + "generate_ecg_demo_clips", + "generate_ppg_demo_clips", "generate_stimulus", + "require_rvq_encoder_quantization_report", "sha256_file", "sync_scorecard_to_deploy", "validate_deploy_package", "write_checksums", "write_json", "write_model_card", + "write_rvq_encoder_quantization_report", "write_scorecard_artifact", ] diff --git a/compressionkit/export/artifact_contract.py b/compressionkit/export/artifact_contract.py index eae737e..b3245cc 100644 --- a/compressionkit/export/artifact_contract.py +++ b/compressionkit/export/artifact_contract.py @@ -24,12 +24,20 @@ class ArtifactFile(StrEnum): DECODER_INT8_HF_TFLITE = "decoder_int8.tflite" DECODER_KERAS = "decoder.keras" DECODER_TFLITE = "decoder.tflite" + DEMO_RECORDINGS = "demo_recordings.npz" + DEMO_RECORDINGS_MANIFEST = "demo_recordings_manifest.json" DENOISER_GAIN_MODEL = "denoiser_gain_model.keras" DENOISER_HEADER = "denoiser_gain_model.h" DENOISER_TFLITE = "denoiser_gain_model.tflite" DENOISER_TRAIN_CONFIG = "denoiser_train_config.json" DEPLOY_MANIFEST = "deploy_manifest.json" ENCODER_HEADER = "encoder.h" + ENCODER_FLOAT32_HEADER = "_encoder_float32.h" + ENCODER_FLOAT32_TFLITE = "encoder_float32.tflite" + ENCODER_FP16_HEADER = "encoder_fp16.h" + ENCODER_FP16_TFLITE = "encoder_fp16.tflite" + ENCODER_INT16X8_HEADER = "encoder_int16x8.h" + ENCODER_INT16X8_TFLITE = "encoder_int16x8.tflite" ENCODER_INT8_HF_TFLITE = "encoder_int8.tflite" ENCODER_KERAS = "encoder.keras" ENCODER_TFLITE = "encoder.tflite" @@ -40,6 +48,7 @@ class ArtifactFile(StrEnum): PRIOR_INT8_TFLITE = "prior_int8.tflite" PRIOR_MANIFEST = "prior_manifest.json" QUALITY_SCORECARD = "quality_scorecard.json" + QUANTIZATION_REPORT = "quantization_report.json" REFERENCE_VECTORS = "reference_vectors.npz" SAMPLE_DATA = "sample_data.npz" SAMPLE_STIMULUS = "sample_stimulus.npz" @@ -59,9 +68,24 @@ class SampleArray(StrEnum): TARGETS = "targets" +class DemoArray(StrEnum): + """Stable arrays in the real-recordings browser-demo NPZ artifact.""" + + SAMPLE_RATE = "sample_rate" + SIGNALS = "signals" + SOURCE_RECORDS = "source_records" + SOURCE_SAMPLE_RATES = "source_sample_rates" + START_SECONDS = "start_seconds" + + RVQ_HF_FILE_RENAMES: tuple[tuple[ArtifactFile, ArtifactFile], ...] = ( (ArtifactFile.ENCODER_TFLITE, ArtifactFile.ENCODER_INT8_HF_TFLITE), (ArtifactFile.ENCODER_HEADER, ArtifactFile.ENCODER_HEADER), + (ArtifactFile.ENCODER_FLOAT32_TFLITE, ArtifactFile.ENCODER_FLOAT32_TFLITE), + (ArtifactFile.ENCODER_FP16_TFLITE, ArtifactFile.ENCODER_FP16_TFLITE), + (ArtifactFile.ENCODER_FP16_HEADER, ArtifactFile.ENCODER_FP16_HEADER), + (ArtifactFile.ENCODER_INT16X8_TFLITE, ArtifactFile.ENCODER_INT16X8_TFLITE), + (ArtifactFile.ENCODER_INT16X8_HEADER, ArtifactFile.ENCODER_INT16X8_HEADER), (ArtifactFile.ENCODER_KERAS, ArtifactFile.ENCODER_KERAS), (ArtifactFile.DECODER_FLOAT32_TFLITE, ArtifactFile.DECODER_FLOAT32_TFLITE), (ArtifactFile.DECODER_TFLITE, ArtifactFile.DECODER_INT8_HF_TFLITE), @@ -73,6 +97,9 @@ class SampleArray(StrEnum): (ArtifactFile.CODEC_SPEC, ArtifactFile.CODEC_SPEC), (ArtifactFile.SAMPLE_DATA, ArtifactFile.SAMPLE_STIMULUS), (ArtifactFile.SAMPLE_STIMULUS, ArtifactFile.SAMPLE_STIMULUS), + (ArtifactFile.DEMO_RECORDINGS, ArtifactFile.DEMO_RECORDINGS), + (ArtifactFile.DEMO_RECORDINGS_MANIFEST, ArtifactFile.DEMO_RECORDINGS_MANIFEST), + (ArtifactFile.QUANTIZATION_REPORT, ArtifactFile.QUANTIZATION_REPORT), (ArtifactFile.DEPLOY_MANIFEST, ArtifactFile.HF_CONFIG), (ArtifactFile.MODEL_CARD, ArtifactFile.MODEL_CARD), (ArtifactFile.PRIOR_INT8_TFLITE, ArtifactFile.PRIOR_INT8_TFLITE), diff --git a/compressionkit/export/demo_ecg.py b/compressionkit/export/demo_ecg.py new file mode 100644 index 0000000..f040bbb --- /dev/null +++ b/compressionkit/export/demo_ecg.py @@ -0,0 +1,299 @@ +"""Select quality-gated, real ECG clips for browser demonstrations. + +The helper draws continuous MLII windows from the MIT-BIH Arrhythmia Database, +then resamples them to the requested demo rate. MIT-BIH is available under the +Open Data Commons Attribution License v1.0; keep the generated manifest and +its source attribution when publishing the resulting clips. +""" + +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass +from pathlib import Path + +import h5py +import numpy as np + +from compressionkit.datasets.ecg import _resample, load_ecg_signal +from compressionkit.evaluation.metrics import compute_ecg_hr_hrv + + +@dataclass(frozen=True) +class EcgDemoQuality: + """Signal-quality measurements recorded for one demo clip.""" + + clipping_fraction: float + dynamic_range: float + heart_rate_bpm: float + num_r_peaks: int + rr_cv: float + standard_deviation: float + + +@dataclass(frozen=True) +class EcgDemoClip: + """One accepted real ECG demo clip and its source provenance.""" + + signal: np.ndarray + quality: EcgDemoQuality + source_record: str + source_sample_rate: int + start_seconds: float + + +@dataclass(frozen=True) +class EcgDemoExport: + """Paths produced when ECG demo clips are written to disk.""" + + csv_paths: list[Path] + manifest_path: Path + + +def assess_ecg_demo_clip(signal: np.ndarray, *, sample_rate: int) -> EcgDemoQuality: + """Validate one single-channel ECG clip for demo suitability. + + The gate rejects non-finite, nearly flat, rail-clipped, implausibly paced, + or highly irregular traces. It favors clean, easy-to-inspect signals for a + waveform-compression demonstration; it is not a clinical quality metric. + + Args: + signal: One-dimensional ECG signal. + sample_rate: Signal sample rate in Hz. + + Returns: + Quality measurements for an accepted clip. + + Raises: + ValueError: If the clip does not meet the demo-quality thresholds. + """ + samples = np.asarray(signal, dtype=np.float32).reshape(-1) + if samples.size < 10 * sample_rate: + raise ValueError("ECG demo clips must be at least 10 seconds long") + if not np.isfinite(samples).all(): + raise ValueError("ECG clip contains non-finite samples") + + standard_deviation = float(np.std(samples)) + dynamic_range = float(np.ptp(samples)) + if standard_deviation < 1e-4 or dynamic_range < 1e-3: + raise ValueError("ECG clip is effectively flat") + rail_tolerance = max(1e-7, dynamic_range * 1e-6) + clipping_fraction = float( + np.mean((samples <= samples.min() + rail_tolerance) | (samples >= samples.max() - rail_tolerance)) + ) + if clipping_fraction > 0.01: + raise ValueError(f"ECG clip appears clipped ({clipping_fraction:.2%} at its rails)") + + hrv = compute_ecg_hr_hrv(samples, sample_rate=sample_rate) + if hrv is None: + raise ValueError("ECG clip has too few detectable R peaks") + + heart_rate_bpm = float(hrv["hr_bpm"]) + peak_locations = np.asarray(hrv["peak_locations"], dtype=np.int64) + rr_intervals = np.diff(peak_locations) / float(sample_rate) + if rr_intervals.size < 3: + raise ValueError("ECG clip has too few R-R intervals") + rr_cv = float(np.std(rr_intervals) / np.mean(rr_intervals)) + if not 45.0 <= heart_rate_bpm <= 110.0: + raise ValueError(f"ECG heart rate {heart_rate_bpm:.1f} BPM is outside the demo range") + if rr_cv > 0.15: + raise ValueError(f"ECG R-R variability {rr_cv:.3f} is too high for a clean demo clip") + + return EcgDemoQuality( + clipping_fraction=clipping_fraction, + dynamic_range=dynamic_range, + heart_rate_bpm=heart_rate_bpm, + num_r_peaks=int(hrv["num_peaks"]), + rr_cv=rr_cv, + standard_deviation=standard_deviation, + ) + + +def _record_metadata(record_path: Path) -> tuple[int, str]: + """Read sample rate and record identifier from one canonical MIT-BIH H5 file.""" + with h5py.File(record_path, "r") as handle: + source_sample_rate = int(handle.attrs["fs"]) + source = str(handle.attrs.get("source", "")) + record_id = str(handle.attrs.get("patient_id", record_path.stem)) + if source != "mitdb": + raise ValueError(f"Expected MIT-BIH record, got source={source!r} in {record_path}") + return source_sample_rate, record_id + + +def generate_ecg_demo_clips( + dataset_dir: str | Path, + *, + num_clips: int = 10, + duration_seconds: float = 30.0, + sample_rate: int = 256, + seed: int = 42, + attempts_per_record: int = 20, +) -> list[EcgDemoClip]: + """Select quality-gated, real MIT-BIH ECG clips for a demo. + + Each accepted clip comes from a distinct record where possible. Windows are + continuous in the source recording before resampling to ``sample_rate``. + + Args: + dataset_dir: Canonical MIT-BIH H5 directory, typically + ``$COMPRESSIONKIT_DATASETS_DIR/mitdb``. + num_clips: Number of clips to select, from 1 through 10. + duration_seconds: Length of each clip in seconds; must be at least 10. + sample_rate: Output sample rate in Hz. + seed: Seed for deterministic record and window selection. + attempts_per_record: Candidate windows to try from each record. + + Returns: + Accepted real ECG demo clips, each sampled at ``sample_rate``. + + Raises: + FileNotFoundError: If no canonical MIT-BIH H5 records are available. + ValueError: If parameters are invalid or too few records pass the + quality gate. + """ + if not 1 <= num_clips <= 10: + raise ValueError("num_clips must be between 1 and 10") + if duration_seconds < 10.0: + raise ValueError("duration_seconds must be at least 10") + if sample_rate <= 0: + raise ValueError("sample_rate must be positive") + if attempts_per_record <= 0: + raise ValueError("attempts_per_record must be positive") + + records = sorted(Path(dataset_dir).glob("*.h5")) + if not records: + raise FileNotFoundError(f"No canonical MIT-BIH H5 records found in {dataset_dir}") + + rng = np.random.default_rng(seed) + selected: list[EcgDemoClip] = [] + for record_path in rng.permutation(records): + source_sample_rate, record_id = _record_metadata(record_path) + source_samples = load_ecg_signal(record_path, lead_index=0) + source_length = round(duration_seconds * source_sample_rate) + if source_samples.size < source_length: + continue + + for _attempt in range(attempts_per_record): + max_start = source_samples.size - source_length + start = int(rng.integers(0, max_start + 1)) + source_window = source_samples[start : start + source_length] + output_window = _resample(source_window, source_sample_rate, sample_rate) + try: + quality = assess_ecg_demo_clip(output_window, sample_rate=sample_rate) + except ValueError: + continue + selected.append( + EcgDemoClip( + signal=output_window, + quality=quality, + source_record=record_id, + source_sample_rate=source_sample_rate, + start_seconds=start / float(source_sample_rate), + ) + ) + break + if len(selected) == num_clips: + return selected + + raise ValueError( + f"Selected only {len(selected)} of {num_clips} quality-gated ECG clips from {dataset_dir}; " + "relax the quality gate or use a larger source dataset." + ) + + +def export_ecg_demo_csvs( + output_dir: str | Path, + *, + dataset_dir: str | Path, + num_clips: int = 10, + duration_seconds: float = 30.0, + sample_rate: int = 256, + seed: int = 42, +) -> EcgDemoExport: + """Write real, quality-gated MIT-BIH ECG clips as CSV files and a manifest. + + Each CSV has ``sample_index``, ``time_s``, and ``ecg`` columns. The JSON + manifest records source attribution, record IDs, source offsets, and the + quality measurements used to accept each clip. + + Args: + output_dir: Destination directory for demo artifacts. + dataset_dir: Canonical MIT-BIH H5 directory. + num_clips: Number of clips to write, from 1 through 10. + duration_seconds: Length of each clip in seconds. + sample_rate: Output sample rate in Hz. + seed: Seed for deterministic record and window selection. + + Returns: + The written CSV paths and quality-manifest path. + """ + destination = Path(output_dir) + destination.mkdir(parents=True, exist_ok=True) + clips = generate_ecg_demo_clips( + dataset_dir, + num_clips=num_clips, + duration_seconds=duration_seconds, + sample_rate=sample_rate, + seed=seed, + ) + + csv_paths: list[Path] = [] + manifest_clips: list[dict[str, object]] = [] + for index, clip in enumerate(clips, start=1): + csv_path = destination / f"ecg_demo_{index:02d}.csv" + sample_indices = np.arange(clip.signal.size, dtype=np.int64) + times = sample_indices / float(sample_rate) + rows = np.column_stack((sample_indices, times, clip.signal)) + np.savetxt( + csv_path, + rows, + delimiter=",", + header="sample_index,time_s,ecg", + comments="", + fmt=["%d", "%.9f", "%.8g"], + ) + csv_paths.append(csv_path) + manifest_clips.append( + { + "file": csv_path.name, + "num_samples": int(clip.signal.size), + "duration_seconds": clip.signal.size / float(sample_rate), + "source_record": clip.source_record, + "source_sample_rate": clip.source_sample_rate, + "start_seconds": clip.start_seconds, + "quality": asdict(clip.quality), + } + ) + + manifest_path = destination / "manifest.json" + manifest_path.write_text( + json.dumps( + { + "modality": "ecg", + "source": { + "dataset": "MIT-BIH Arrhythmia Database v1.0.0", + "license": "ODC-By-1.0", + "citation": "Moody GB, Mark RG. The impact of the MIT-BIH Arrhythmia Database. " + "IEEE Eng Med Biol. 2001;20(3):45-50.", + }, + "sample_rate": sample_rate, + "num_clips": len(clips), + "duration_seconds": duration_seconds, + "seed": seed, + "clips": manifest_clips, + }, + indent=2, + ) + + "\n" + ) + return EcgDemoExport(csv_paths=csv_paths, manifest_path=manifest_path) + + +__all__ = [ + "EcgDemoClip", + "EcgDemoExport", + "EcgDemoQuality", + "assess_ecg_demo_clip", + "export_ecg_demo_csvs", + "generate_ecg_demo_clips", +] diff --git a/compressionkit/export/demo_ppg.py b/compressionkit/export/demo_ppg.py new file mode 100644 index 0000000..7d44b29 --- /dev/null +++ b/compressionkit/export/demo_ppg.py @@ -0,0 +1,205 @@ +"""Select quality-gated, real PPG clips for browser demonstrations. + +The helper reads continuous PLETH windows from the BIDMC PPG and Respiration +Dataset and resamples them to the target demo rate. BIDMC is available under +the Open Data Commons Attribution License v1.0; retain the generated manifest +when distributing these clips. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +import h5py +import numpy as np + +from compressionkit.datasets.ppg_h5 import _resample +from compressionkit.evaluation.metrics import compute_ppg_physiokit_metrics +from compressionkit.preprocessing.curation import ppg_pulse_snr_db + + +@dataclass(frozen=True) +class PpgDemoQuality: + """Signal-quality measurements recorded for one demo clip.""" + + clipping_fraction: float + dynamic_range: float + heart_rate_bpm: float + heart_rate_qos: float + num_peaks: int + pulse_snr_db: float + rr_cv: float + standard_deviation: float + + +@dataclass(frozen=True) +class PpgDemoClip: + """One accepted real PPG demo clip and its source provenance.""" + + signal: np.ndarray + quality: PpgDemoQuality + source_record: str + source_sample_rate: int + start_seconds: float + + +def assess_ppg_demo_clip(signal: np.ndarray, *, sample_rate: int) -> PpgDemoQuality: + """Validate one single-channel PPG clip for demo suitability. + + This gate favors readable, continuous waveforms with a stable pulse rhythm. + It is only a demonstration-quality filter and is not a clinical SQI. + + Args: + signal: One-dimensional PPG signal. + sample_rate: Signal sample rate in Hz. + + Returns: + Quality measurements for an accepted clip. + + Raises: + ValueError: If the clip does not meet the demo-quality thresholds. + """ + samples = np.asarray(signal, dtype=np.float32).reshape(-1) + if samples.size < 10 * sample_rate: + raise ValueError("PPG demo clips must be at least 10 seconds long") + if not np.isfinite(samples).all(): + raise ValueError("PPG clip contains non-finite samples") + + standard_deviation = float(np.std(samples)) + dynamic_range = float(np.ptp(samples)) + if standard_deviation < 1e-4 or dynamic_range < 1e-3: + raise ValueError("PPG clip is effectively flat") + rail_tolerance = max(1e-7, dynamic_range * 1e-6) + clipping_fraction = float( + np.mean((samples <= samples.min() + rail_tolerance) | (samples >= samples.max() - rail_tolerance)) + ) + if clipping_fraction > 0.01: + raise ValueError(f"PPG clip appears clipped ({clipping_fraction:.2%} at its rails)") + + metrics = compute_ppg_physiokit_metrics( + samples, + sample_rate=sample_rate, + low_hz=0.5, + high_hz=min(8.0, sample_rate / 2.5), + order=3, + min_peaks=5, + ) + if metrics is None: + raise ValueError("PPG clip has too few detectable pulse peaks") + heart_rate_bpm = float(metrics["hr_bpm"]) + heart_rate_qos = float(metrics["hr_qos"]) + peak_locations = np.asarray(metrics["peak_locations"], dtype=np.int64) + rr_intervals = np.diff(peak_locations) / float(sample_rate) + if rr_intervals.size < 3: + raise ValueError("PPG clip has too few pulse intervals") + rr_cv = float(np.std(rr_intervals) / np.mean(rr_intervals)) + pulse_snr_db = float(ppg_pulse_snr_db(samples, sample_rate)) + if not 45.0 <= heart_rate_bpm <= 120.0: + raise ValueError(f"PPG heart rate {heart_rate_bpm:.1f} BPM is outside the demo range") + if heart_rate_qos < 0.8: + raise ValueError(f"PPG heart-rate quality {heart_rate_qos:.2f} is too low for a demo") + if rr_cv > 0.2: + raise ValueError(f"PPG pulse-interval variability {rr_cv:.3f} is too high for a clean demo clip") + if not np.isfinite(pulse_snr_db) or pulse_snr_db < 8.0: + raise ValueError(f"PPG pulse-template SNR {pulse_snr_db:.1f} dB is too low for a demo") + + return PpgDemoQuality( + clipping_fraction=clipping_fraction, + dynamic_range=dynamic_range, + heart_rate_bpm=heart_rate_bpm, + heart_rate_qos=heart_rate_qos, + num_peaks=int(metrics["num_peaks"]), + pulse_snr_db=pulse_snr_db, + rr_cv=rr_cv, + standard_deviation=standard_deviation, + ) + + +def _record_metadata(record_path: Path) -> tuple[int, str]: + """Read source metadata from one canonical BIDMC H5 file.""" + with h5py.File(record_path, "r") as handle: + source_sample_rate = int(handle.attrs["fs"]) + source = str(handle.attrs.get("source", "")) + record_id = str(handle.attrs.get("patient_id", record_path.stem)) + if source != "bidmc": + raise ValueError(f"Expected BIDMC record, got source={source!r} in {record_path}") + return source_sample_rate, record_id + + +def generate_ppg_demo_clips( + dataset_dir: str | Path, + *, + num_clips: int = 10, + duration_seconds: float = 30.0, + sample_rate: int = 64, + seed: int = 42, + attempts_per_record: int = 20, +) -> list[PpgDemoClip]: + """Select quality-gated, real BIDMC PPG clips for a demo. + + Each accepted clip comes from a distinct BIDMC record where possible. + Windows are continuous in the source recording before resampling. + + Args: + dataset_dir: Canonical BIDMC H5 directory. + num_clips: Number of clips to select, from 1 through 10. + duration_seconds: Length of each clip in seconds; must be at least 10. + sample_rate: Output sample rate in Hz. + seed: Seed for deterministic record and window selection. + attempts_per_record: Candidate windows to try from each record. + + Returns: + Accepted real PPG clips, each sampled at ``sample_rate``. + """ + if not 1 <= num_clips <= 10: + raise ValueError("num_clips must be between 1 and 10") + if duration_seconds < 10.0: + raise ValueError("duration_seconds must be at least 10") + if sample_rate <= 0: + raise ValueError("sample_rate must be positive") + if attempts_per_record <= 0: + raise ValueError("attempts_per_record must be positive") + + records = sorted(Path(dataset_dir).glob("*.h5")) + if not records: + raise FileNotFoundError(f"No canonical BIDMC H5 records found in {dataset_dir}") + + rng = np.random.default_rng(seed) + selected: list[PpgDemoClip] = [] + for record_path in rng.permutation(records): + source_sample_rate, record_id = _record_metadata(record_path) + with h5py.File(record_path, "r") as handle: + source_samples = handle["data"][0].astype(np.float32, copy=False) + source_length = round(duration_seconds * source_sample_rate) + if source_samples.size < source_length: + continue + + for _attempt in range(attempts_per_record): + max_start = source_samples.size - source_length + start = int(rng.integers(0, max_start + 1)) + output_window = _resample(source_samples[start : start + source_length], source_sample_rate, sample_rate) + try: + quality = assess_ppg_demo_clip(output_window, sample_rate=sample_rate) + except ValueError: + continue + selected.append( + PpgDemoClip( + signal=output_window, + quality=quality, + source_record=record_id, + source_sample_rate=source_sample_rate, + start_seconds=start / float(source_sample_rate), + ) + ) + break + if len(selected) == num_clips: + return selected + + raise ValueError( + f"Selected only {len(selected)} of {num_clips} quality-gated PPG clips from {dataset_dir}; " + "relax the quality gate or use a larger source dataset." + ) + + +__all__ = ["PpgDemoClip", "PpgDemoQuality", "assess_ppg_demo_clip", "generate_ppg_demo_clips"] diff --git a/compressionkit/export/demo_recordings.py b/compressionkit/export/demo_recordings.py new file mode 100644 index 0000000..ff11d77 --- /dev/null +++ b/compressionkit/export/demo_recordings.py @@ -0,0 +1,136 @@ +"""Package curated real recordings as stable deployment artifacts.""" + +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Protocol + +import numpy as np + +from compressionkit.export.artifact_contract import ArtifactFile, DemoArray +from compressionkit.export.release import write_checksums, write_json + + +class _DemoClip(Protocol): + """Minimal common interface implemented by ECG and PPG demo clips.""" + + signal: np.ndarray + quality: object + source_record: str + source_sample_rate: int + start_seconds: float + + +@dataclass(frozen=True) +class DemoRecordingsExport: + """Paths written for a portable browser-demo recording bundle.""" + + manifest_path: Path + recordings_path: Path + + +def export_demo_recordings( + output_dir: str | Path, + *, + clips: list[_DemoClip], + modality: str, + sample_rate: int, + source: dict[str, str], + seed: int, +) -> DemoRecordingsExport: + """Write a compact real-recordings NPZ and provenance manifest. + + The recordings are intentionally raw source waveforms after resampling, so + a web application can apply the same framing and normalization policy as + its selected codec. Every clip must share output length and sample rate. + + Args: + output_dir: Destination deploy directory or standalone output directory. + clips: Quality-gated real demo clips from one modality. + modality: ``"ecg"`` or ``"ppg"``. + sample_rate: Output signal rate in Hz. + source: Dataset attribution and license fields. + seed: Deterministic source-selection seed for the manifest. + + Returns: + Paths of the generated NPZ and JSON manifest. + """ + if modality not in {"ecg", "ppg"}: + raise ValueError(f"Unsupported demo modality {modality!r}") + if not clips: + raise ValueError("At least one demo clip is required") + if sample_rate <= 0: + raise ValueError("sample_rate must be positive") + + destination = Path(output_dir) + destination.mkdir(parents=True, exist_ok=True) + signals = np.stack([np.asarray(clip.signal, dtype=np.float32).reshape(-1) for clip in clips]) + if signals.ndim != 2: + raise ValueError("Demo clips must be one-dimensional and have equal lengths") + + recordings_path = destination / ArtifactFile.DEMO_RECORDINGS + np.savez_compressed( + recordings_path, + **{ + DemoArray.SIGNALS: signals, + DemoArray.SAMPLE_RATE: np.asarray(sample_rate, dtype=np.int32), + DemoArray.SOURCE_RECORDS: np.asarray([clip.source_record for clip in clips]), + DemoArray.SOURCE_SAMPLE_RATES: np.asarray([clip.source_sample_rate for clip in clips], dtype=np.int32), + DemoArray.START_SECONDS: np.asarray([clip.start_seconds for clip in clips], dtype=np.float32), + }, + ) + manifest = { + "format_version": 1, + "modality": modality, + "sample_rate": sample_rate, + "num_recordings": len(clips), + "samples_per_recording": int(signals.shape[1]), + "duration_seconds": signals.shape[1] / float(sample_rate), + "signal_representation": "raw source waveform resampled to sample_rate", + "framing_note": "Apply the selected codec's framing and normalization policy before inference.", + "source": source, + "seed": seed, + "recordings": [ + { + "index": index, + "source_record": clip.source_record, + "source_sample_rate": clip.source_sample_rate, + "start_seconds": clip.start_seconds, + "quality": asdict(clip.quality), + } + for index, clip in enumerate(clips) + ], + } + manifest_path = destination / ArtifactFile.DEMO_RECORDINGS_MANIFEST + manifest_path.write_text(json.dumps(manifest, indent=2) + "\n") + return DemoRecordingsExport(manifest_path=manifest_path, recordings_path=recordings_path) + + +def attach_demo_recordings_to_deploy( + deploy_dir: str | Path, + bundle: DemoRecordingsExport, +) -> None: + """Register an existing demo bundle in a deploy manifest and checksums.""" + destination = Path(deploy_dir) + manifest_path = destination / ArtifactFile.DEPLOY_MANIFEST + if bundle.recordings_path.parent != destination or bundle.manifest_path.parent != destination: + raise ValueError("Demo bundle must be written directly inside the deploy directory") + manifest = json.loads(manifest_path.read_text()) + recordings = json.loads(bundle.manifest_path.read_text()) + manifest["demo_recordings"] = { + "npz": bundle.recordings_path.name, + "manifest": bundle.manifest_path.name, + "num_recordings": recordings["num_recordings"], + "sample_rate": recordings["sample_rate"], + "duration_seconds": recordings["duration_seconds"], + } + artifacts = manifest.setdefault("artifacts", {}) + artifacts["demo_recordings"] = bundle.recordings_path.name + artifacts["demo_recordings_manifest"] = bundle.manifest_path.name + write_json(manifest_path, manifest) + write_checksums(destination) + + +__all__ = ["DemoRecordingsExport", "attach_demo_recordings_to_deploy", "export_demo_recordings"] diff --git a/compressionkit/export/deploy.py b/compressionkit/export/deploy.py index f08a492..d8510de 100644 --- a/compressionkit/export/deploy.py +++ b/compressionkit/export/deploy.py @@ -1,7 +1,7 @@ """Unified deployment export for RVQ autoencoder models. Exports all components needed for deployment: - 1. Encoder — INT8 TFLite + C header (on-device) + 1. Encoder — INT8 TFLite + C header and FP32 LiteRT (edge/browser) 2. Codebook tables — C header + NumPy archive (on-device) 3. Decoder — Keras model + optional TFLite exports (server / on-device) 4. Model card JSON with metadata and scorecard summary @@ -40,6 +40,11 @@ class DeploymentArtifacts: output_dir: Path encoder_tflite: Path = field(default_factory=Path) encoder_header: Path = field(default_factory=Path) + encoder_float32_tflite: Path = field(default_factory=Path) + encoder_fp16_tflite: Path = field(default_factory=Path) + encoder_fp16_header: Path = field(default_factory=Path) + encoder_int16x8_tflite: Path = field(default_factory=Path) + encoder_int16x8_header: Path = field(default_factory=Path) encoder_keras: Path = field(default_factory=Path) decoder_keras: Path = field(default_factory=Path) decoder_float32_tflite: Path = field(default_factory=Path) @@ -49,6 +54,7 @@ class DeploymentArtifacts: codebook_header: Path = field(default_factory=Path) sample_data_npz: Path = field(default_factory=Path) reference_vectors: Path = field(default_factory=Path) + quantization_report: Path = field(default_factory=Path) model_card: Path = field(default_factory=Path) scorecard: Path = field(default_factory=Path) codec_spec: Path = field(default_factory=Path) @@ -95,7 +101,9 @@ def _render_readme( * `deploy_manifest.json` — top-level package manifest. * `codec_spec.json` — canonical runtime hydration contract. -* `encoder.tflite`, `encoder.h` — edge encoder artifacts. +* `encoder_float32.tflite` — default FP32 encoder for browser and host LiteRT runtimes. +* `encoder.tflite`, `encoder_fp16.tflite`, `encoder_int16x8.tflite` — INT8, FP16, + and INT16x8 edge encoder variants (each with a C header). * `encoder.keras` — float32 Python reference encoder (training/inspection use). * `codebook.npz`, `codebook.h` — RVQ codebook tables. * `decoder.keras` and optional decoder TFLite files — reconstruction artifacts. @@ -152,6 +160,7 @@ def export_for_deployment( rvq_weights: list[np.ndarray], *, rep_dataset: np.ndarray, + quantization_validation_dataset: np.ndarray | None = None, output_dir: str | Path, sample_inputs: np.ndarray | None = None, sample_targets: np.ndarray | None = None, @@ -176,6 +185,8 @@ def export_for_deployment( decoder: Trained Keras decoder model. rvq_weights: Weight list from ``rvq.get_weights()``. rep_dataset: Representative input dataset for TFLite calibration. + quantization_validation_dataset: Disjoint preprocessed frames for + post-export INT8-vs-float32 parity validation. Required for INT8. output_dir: Directory to write all deployment artifacts. sample_inputs: Optional sample input frames for validation. sample_targets: Optional sample target frames for validation. @@ -196,6 +207,8 @@ def export_for_deployment( ``DeploymentArtifacts`` with paths to all exported files. """ output_dir = Path(output_dir) + if quantization == "INT8" and quantization_validation_dataset is None: + raise ValueError("INT8 export requires a disjoint quantization_validation_dataset") output_dir.mkdir(parents=True, exist_ok=True) artifacts = DeploymentArtifacts(output_dir=output_dir) @@ -212,7 +225,50 @@ def export_for_deployment( artifacts.encoder_tflite = enc_tflite artifacts.encoder_header = enc_header - # 1b. Encoder (.keras -- float32 reference alongside the quantized + # 1b. Float32 encoder LiteRT (browser/host runtime without Keras). + logger.info("Exporting float32 encoder TFLite...") + enc_f32_tflite, _ = export_encoder_tflite( + encoder, + rep_dataset=rep_dataset, + output_dir=output_dir, + tflite_name="encoder_float32.tflite", + header_name="_encoder_float32.h", + c_array_name="encoder_float32", + quantization="FP32", + io_type="float32", + ) + artifacts.encoder_float32_tflite = enc_f32_tflite + + # 1c. Additional MCU encoder precisions. ``encoder.tflite`` remains the + # backwards-compatible INT8 filename; explicit names let an application + # select its HeliaAOT-supported precision without changing codec payloads. + logger.info("Exporting FP16 and INT16x8 encoder variants...") + enc_fp16_tflite, enc_fp16_header = export_encoder_tflite( + encoder, + rep_dataset=rep_dataset, + output_dir=output_dir, + tflite_name="encoder_fp16.tflite", + header_name="encoder_fp16.h", + c_array_name="encoder_fp16", + quantization="FP16", + io_type="float32", + ) + enc_int16x8_tflite, enc_int16x8_header = export_encoder_tflite( + encoder, + rep_dataset=rep_dataset, + output_dir=output_dir, + tflite_name="encoder_int16x8.tflite", + header_name="encoder_int16x8.h", + c_array_name="encoder_int16x8", + quantization="INT16X8", + io_type="float32", + ) + artifacts.encoder_fp16_tflite = enc_fp16_tflite + artifacts.encoder_fp16_header = enc_fp16_header + artifacts.encoder_int16x8_tflite = enc_int16x8_tflite + artifacts.encoder_int16x8_header = enc_int16x8_header + + # 1d. Encoder (.keras -- float32 reference alongside the quantized # encoder.tflite used on-device; mirrors the decoder.keras reference below). logger.info("Exporting encoder as .keras...") encoder_keras_path = output_dir / "encoder.keras" @@ -320,12 +376,27 @@ def export_for_deployment( "package_version": model_version, "input_contract": { "encoder_input_shape": list(encoder.input_shape), - "encoder_input_dtype": io_type, + "encoder_input_dtype": "float32", }, "preprocessing_contract": preprocessing_contract, "encoder": { - "tflite": enc_tflite.name, + "tflite": enc_f32_tflite.name, "header": enc_header.name, + "default_variant": "float32", + "default_tflite": enc_f32_tflite.name, + "float32_tflite": enc_f32_tflite.name, + "fp16_tflite": enc_fp16_tflite.name, + "fp16_header": enc_fp16_header.name, + "int16x8_tflite": enc_int16x8_tflite.name, + "int16x8_header": enc_int16x8_header.name, + "int8_tflite": enc_tflite.name, + "int8_header": enc_header.name, + "variants": { + "float32": {"tflite": enc_f32_tflite.name, "input_dtype": "float32"}, + "fp16": {"tflite": enc_fp16_tflite.name, "input_dtype": "float32"}, + "int16x8": {"tflite": enc_int16x8_tflite.name, "input_dtype": "float32"}, + "int8": {"tflite": enc_tflite.name, "input_dtype": io_type}, + }, "keras": encoder_keras_path.name, "input_shape": list(encoder.input_shape), "output_shape": list(encoder.output_shape), @@ -398,8 +469,17 @@ def export_for_deployment( "io_type": io_type, "artifacts": artifacts.as_dict(), "encoder": { - "tflite": enc_tflite.name, + "tflite": enc_f32_tflite.name, "header": enc_header.name, + "default_variant": "float32", + "default_tflite": enc_f32_tflite.name, + "float32_tflite": enc_f32_tflite.name, + "fp16_tflite": enc_fp16_tflite.name, + "fp16_header": enc_fp16_header.name, + "int16x8_tflite": enc_int16x8_tflite.name, + "int16x8_header": enc_int16x8_header.name, + "int8_tflite": enc_tflite.name, + "int8_header": enc_header.name, "keras": encoder_keras_path.name, "input_shape": list(encoder.input_shape), "output_shape": list(encoder.output_shape), @@ -460,6 +540,21 @@ def export_for_deployment( ) artifacts.reference_vectors = Path() + # 7b. Compare the deployed INT8 encoder against its float32 companion on + # a disjoint, real preprocessed validation partition. + if quantization == "INT8": + from compressionkit.export.quantization import ( + evaluate_rvq_encoder_quantization, + write_rvq_encoder_quantization_report, + ) + + report = evaluate_rvq_encoder_quantization(output_dir, quantization_validation_dataset, max_frames=None) + artifacts.quantization_report = write_rvq_encoder_quantization_report( + output_dir, + report, + calibration_frames=int(np.asarray(rep_dataset).shape[0]), + ) + # 8. File-integrity manifest (written last so it can include everything else). artifacts.checksums = write_checksums(output_dir) diff --git a/compressionkit/export/family_registry.py b/compressionkit/export/family_registry.py index 10ea19d..5957816 100644 --- a/compressionkit/export/family_registry.py +++ b/compressionkit/export/family_registry.py @@ -97,12 +97,19 @@ class CodecFamilySpec: loader=_load_rvq, required_artifacts=( "encoder.tflite", + "encoder_float32.tflite", "encoder.keras", "decoder.keras", "codebook.npz", "codebook.h", ), - release_extras=("model_card.json", "README.md", "scorecard.json", "reference_vectors.npz", "sample_data.npz"), + release_extras=( + "model_card.json", + "README.md", + "scorecard.json", + "reference_vectors.npz", + "sample_data.npz", + ), hf_file_renames=tuple(RVQ_HF_FILE_RENAMES), has_c_sources=False, has_trained_weights=True, diff --git a/compressionkit/export/model_card.py b/compressionkit/export/model_card.py index d1180a8..84886b8 100644 --- a/compressionkit/export/model_card.py +++ b/compressionkit/export/model_card.py @@ -194,6 +194,23 @@ def generate_model_card( lines.append("") _add_scorecard_section(lines, scorecard) + precision_report = model_card_info.get("encoder_precision_report") + if isinstance(precision_report, dict): + lines.append("## Encoder Precision Parity") + lines.append("") + lines.append("Difference from FP32 reconstruction on a disjoint real-data holdout; lower is better.") + lines.append("") + lines.append("| Encoder | P90 PRD | Worst PRD | Status |") + lines.append("|---|---:|---:|---|") + for name, metrics in precision_report.items(): + if not isinstance(metrics, dict): + continue + p90 = float(metrics.get("reconstruction_prd_percent_p90", 0.0)) + maximum = float(metrics.get("reconstruction_prd_percent_max", 0.0)) + status = "recommended" if p90 <= 10.0 else "higher error; use knowingly" + lines.append(f"| {name} | {p90:.2f}% | {maximum:.2f}% | {status} |") + lines.append("") + # Usage lines.append("## Usage") lines.append("") @@ -224,6 +241,9 @@ def generate_model_card( lines.append("| File | Description |") lines.append("|------|-------------|") lines.append("| `encoder_int8.tflite` | INT8 quantized encoder (on-device) |") + lines.append("| `encoder_float32.tflite` | Float32 encoder for browser/server runtimes |") + lines.append("| `encoder_fp16.tflite` | FP16 encoder variant for supported edge runtimes |") + lines.append("| `encoder_int16x8.tflite` | INT16x8 encoder variant for supported edge runtimes |") lines.append("| `encoder.h` | C header for encoder |") lines.append("| `encoder.keras` | Float32 Python reference encoder (training/inspection use) |") lines.append("| `decoder_float32.tflite` | Float32 decoder (server-side evaluation) |") @@ -233,6 +253,9 @@ def generate_model_card( lines.append("| `codebook.h` | C header for codebook |") lines.append("| `config.json` | Deployment manifest |") lines.append("| `sample_stimulus.npz` | Synthetic test data |") + if "demo_recordings" in manifest: + lines.append("| `demo_recordings.npz` | 10 real, quality-gated browser-demo recordings |") + lines.append("| `demo_recordings_manifest.json` | Recording provenance, license, and quality metadata |") lines.append("| `quality_scorecard.json` | Full evaluation metrics |") lines.append("") @@ -240,6 +263,23 @@ def generate_model_card( lines.append("## Dataset & License") lines.append("") lines.append(_dataset_provenance_note(model_card_info.get("dataset_sources"), modality, context="Training")) + demo_recordings = manifest.get("demo_recordings", {}) + if isinstance(demo_recordings, dict): + demo_manifest_path = deploy_dir / str(demo_recordings.get("manifest", "")) + if demo_manifest_path.is_file(): + demo_manifest = json.loads(demo_manifest_path.read_text()) + source = demo_manifest.get("source", {}) + if isinstance(source, dict): + dataset = source.get("dataset", "the source dataset") + license_name = source.get("license", "its original license") + license_url = source.get("license_url") + source_url = source.get("url") + reference = f" [{license_url}]({license_url})" if license_url else "" + source_link = f" ([source]({source_url}))" if source_url else "" + lines.append( + f"Demo recordings: {dataset}{source_link}; real, quality-gated excerpts are released under " + f"{license_name}.{reference}" + ) lines.append("") if license_id == "other": lines.append( diff --git a/compressionkit/export/quantization.py b/compressionkit/export/quantization.py new file mode 100644 index 0000000..b4a340a --- /dev/null +++ b/compressionkit/export/quantization.py @@ -0,0 +1,242 @@ +"""Numerical parity checks for RVQ encoder precision exports.""" + +from __future__ import annotations + +import json +from dataclasses import asdict, dataclass +from pathlib import Path + +import numpy as np + +from compressionkit.export.artifact_contract import ArtifactFile +from compressionkit.export.release import write_checksums, write_json +from compressionkit.runtime._litert import Interpreter, dequantize, quantize +from compressionkit.runtime.codec import RVQCodec + +MAX_INPUT_SATURATION_FRACTION = 0.01 +MAX_RECONSTRUCTION_PRD_PERCENT_P90 = 10.0 +TAIL_WARNING_RECONSTRUCTION_PRD_PERCENT = 15.0 + + +@dataclass(frozen=True) +class RvqEncoderQuantizationReport: + """Parity and saturation measurements for one RVQ encoder export.""" + + code_index_match_fraction_median: float + encoder_input_saturation_fraction_max: float + encoder_input_saturation_fraction_p90: float + frames_checked: int + latent_prd_percent_p90: float + reconstruction_prd_percent_max: float + reconstruction_prd_percent_p90: float + + def passed(self) -> bool: + """Return whether this export meets the population-quality release gate. + + Input saturation is the hard publication criterion. P90 and + worst-frame PRD are retained as explicit quality recommendations so + customers can select an encoder precision with full visibility. + """ + return ( + self.encoder_input_saturation_fraction_max <= MAX_INPUT_SATURATION_FRACTION + ) + + +def _prd_percent(actual: np.ndarray, expected: np.ndarray) -> float: + """Return RMS percent difference, with a safe zero-reference fallback.""" + err = np.asarray(actual, dtype=np.float64) - np.asarray(expected, dtype=np.float64) + denominator = float(np.sum(np.asarray(expected, dtype=np.float64) ** 2)) + if denominator <= 1e-12: + return 0.0 if np.allclose(actual, expected) else float("inf") + return 100.0 * float(np.sqrt(np.sum(err**2) / denominator)) + + +def evaluate_rvq_encoder_quantization( + deploy_dir: str | Path, + frames: np.ndarray, + *, + max_frames: int | None = 128, + candidate_encoder_name: str = ArtifactFile.ENCODER_TFLITE, +) -> RvqEncoderQuantizationReport: + """Compare one RVQ encoder variant with the float32 reference encoder. + + Args: + deploy_dir: Complete RVQ deploy directory containing both encoder variants. + frames: Input frames after the exact training-time preprocessing. + max_frames: Maximum number of frames to evaluate, or ``None`` for all. + candidate_encoder_name: Candidate encoder filename relative to + ``deploy_dir``. Defaults to the standard INT8 encoder. + + Returns: + Candidate precision parity and input-saturation measurements. + + Raises: + ValueError: If no frames are supplied or their shape is incompatible. + """ + samples = np.asarray(frames, dtype=np.float32) + if samples.ndim != 4 or samples.shape[0] == 0: + raise ValueError(f"Expected non-empty (N, 1, T, C) frames, got {samples.shape}") + if max_frames is not None: + samples = samples[:max_frames] + root = Path(deploy_dir) + candidate_encoder = Interpreter(model_path=str(root / candidate_encoder_name)) + float_encoder = Interpreter(model_path=str(root / ArtifactFile.ENCODER_FLOAT32_TFLITE)) + candidate_encoder.allocate_tensors() + float_encoder.allocate_tensors() + candidate_input, candidate_output = candidate_encoder.get_input_details()[0], candidate_encoder.get_output_details()[0] + float_input, float_output = float_encoder.get_input_details()[0], float_encoder.get_output_details()[0] + codec = RVQCodec(root) + + input_saturation: list[float] = [] + latent_prd: list[float] = [] + reconstruction_prd: list[float] = [] + index_match: list[float] = [] + candidate_is_integer = np.issubdtype(candidate_input["dtype"], np.integer) + limits = np.iinfo(candidate_input["dtype"]) if candidate_is_integer else None + for frame in samples: + batch = frame[np.newaxis, ...] + candidate_input_data = quantize(batch, candidate_input) + if limits is None: + input_saturation.append(0.0) + else: + input_saturation.append( + float(np.mean((candidate_input_data == limits.min) | (candidate_input_data == limits.max))) + ) + candidate_encoder.set_tensor(candidate_input["index"], candidate_input_data) + candidate_encoder.invoke() + candidate_latent = dequantize(candidate_encoder.get_tensor(candidate_output["index"]), candidate_output) + float_encoder.set_tensor(float_input["index"], batch.astype(float_input["dtype"])) + float_encoder.invoke() + float_latent = float_encoder.get_tensor(float_output["index"]).astype(np.float32) + candidate_indices = codec.quantize_latent(candidate_latent) + float_indices = codec.quantize_latent(float_latent) + candidate_reconstruction = codec.decode_latent(codec.dequantize_indices(candidate_indices)) + float_reconstruction = codec.decode_latent(codec.dequantize_indices(float_indices)) + latent_prd.append(_prd_percent(candidate_latent, float_latent)) + reconstruction_prd.append(_prd_percent(candidate_reconstruction, float_reconstruction)) + index_match.append(float(np.mean(candidate_indices == float_indices))) + + return RvqEncoderQuantizationReport( + frames_checked=int(samples.shape[0]), + encoder_input_saturation_fraction_max=float(np.max(input_saturation)), + encoder_input_saturation_fraction_p90=float(np.quantile(input_saturation, 0.9)), + latent_prd_percent_p90=float(np.quantile(latent_prd, 0.9)), + reconstruction_prd_percent_p90=float(np.quantile(reconstruction_prd, 0.9)), + reconstruction_prd_percent_max=float(np.max(reconstruction_prd)), + code_index_match_fraction_median=float(np.median(index_match)), + ) + + +def write_rvq_encoder_quantization_report( + deploy_dir: str | Path, + report: RvqEncoderQuantizationReport, + *, + calibration_frames: int, +) -> Path: + """Persist a passed quantization report and register it in the deploy manifest. + + Raises: + ValueError: If the report exceeds the release thresholds. + """ + if not report.passed(): + raise ValueError( + "INT8 encoder quantization parity failed: " + f"saturation max={report.encoder_input_saturation_fraction_max:.2%}, " + f"reconstruction PRD p90={report.reconstruction_prd_percent_p90:.2f}%" + ) + root = Path(deploy_dir) + report_path = root / ArtifactFile.QUANTIZATION_REPORT + payload = { + "format_version": 1, + "passed": True, + "thresholds": { + "max_input_saturation_fraction": MAX_INPUT_SATURATION_FRACTION, + "p90_reconstruction_prd_percent": MAX_RECONSTRUCTION_PRD_PERCENT_P90, + "tail_warning_reconstruction_prd_percent": TAIL_WARNING_RECONSTRUCTION_PRD_PERCENT, + }, + "warnings": [ + *( + [ + "P90 reconstruction PRD exceeds the recommended threshold: " + f"{report.reconstruction_prd_percent_p90:.2f}%" + ] + if report.reconstruction_prd_percent_p90 > MAX_RECONSTRUCTION_PRD_PERCENT_P90 + else [] + ), + *( + [ + "worst-frame reconstruction PRD exceeds the tail-warning threshold: " + f"{report.reconstruction_prd_percent_max:.2f}%" + ] + if report.reconstruction_prd_percent_max > TAIL_WARNING_RECONSTRUCTION_PRD_PERCENT + else [] + ), + ], + "metrics": asdict(report), + "partitions": { + "calibration_frames": calibration_frames, + "validation_frames": report.frames_checked, + "disjoint": True, + }, + } + write_json(report_path, payload) + manifest_path = root / ArtifactFile.DEPLOY_MANIFEST + manifest = json.loads(manifest_path.read_text()) + manifest["quantization_validation"] = { + "report": report_path.name, + "passed": True, + "calibration_frames": calibration_frames, + "validation_frames": report.frames_checked, + "disjoint": True, + } + manifest.setdefault("artifacts", {})["quantization_report"] = report_path.name + write_json(manifest_path, manifest) + write_checksums(root) + return report_path + + +def require_rvq_encoder_quantization_report(deploy_dir: str | Path) -> None: + """Reject an INT8 RVQ package without a passed quantization report. + + This lightweight guard is used by the Hugging Face publisher. It does not + replay models; full artifact and runtime verification remains the job of + :func:`compressionkit.export.validate.validate_deploy_package`. + + Args: + deploy_dir: RVQ deploy directory containing ``deploy_manifest.json``. + + Raises: + ValueError: If an INT8 RVQ package has no valid passed report. + """ + root = Path(deploy_dir) + manifest_path = root / ArtifactFile.DEPLOY_MANIFEST + manifest = json.loads(manifest_path.read_text()) + if manifest.get("family", "rvq") != "rvq" or manifest.get("quantization") != "INT8": + return + validation = manifest.get("quantization_validation") + if not isinstance(validation, dict) or validation.get("passed") is not True: + raise ValueError("INT8 RVQ release is missing a passed quantization_validation entry") + report_name = validation.get("report", ArtifactFile.QUANTIZATION_REPORT) + report_path = root / str(report_name) + if not report_path.is_file(): + raise ValueError(f"INT8 RVQ release is missing quantization report: {report_name}") + report = json.loads(report_path.read_text()) + if report.get("passed") is not True: + raise ValueError("INT8 RVQ release quantization report is not passed") + metrics = report.get("metrics") + if not isinstance(metrics, dict): + raise ValueError("INT8 RVQ release quantization report is missing metrics") + try: + measured = RvqEncoderQuantizationReport(**metrics) + except TypeError as exc: + raise ValueError(f"INT8 RVQ release quantization report has invalid metrics: {exc}") from exc + if not measured.passed(): + raise ValueError("INT8 RVQ release quantization report exceeds release thresholds") + + +__all__ = [ + "RvqEncoderQuantizationReport", + "evaluate_rvq_encoder_quantization", + "require_rvq_encoder_quantization_report", + "write_rvq_encoder_quantization_report", +] diff --git a/compressionkit/export/validate.py b/compressionkit/export/validate.py index 12c224f..cac593f 100644 --- a/compressionkit/export/validate.py +++ b/compressionkit/export/validate.py @@ -226,6 +226,67 @@ def validate_deploy_package( for rel in family_spec.release_extras if family_spec is not None else default_release_extras: _check_file(root, rel, checked_files, errors, warnings, required=strict_release, label="release artifact") + # Real browser-demo recordings are attached in a separate, source-data + # release step. They are optional for a clean model export, but once the + # manifest advertises them both artifacts are part of the package contract. + demo_recordings = manifest.get("demo_recordings") + if demo_recordings is not None: + if not isinstance(demo_recordings, dict): + errors.append("demo_recordings manifest entry must be an object") + else: + _check_file( + root, + str(demo_recordings.get("npz", "demo_recordings.npz")), + checked_files, + errors, + warnings, + required=True, + label="demo recordings artifact", + ) + _check_file( + root, + str(demo_recordings.get("manifest", "demo_recordings_manifest.json")), + checked_files, + errors, + warnings, + required=True, + label="demo recordings artifact", + ) + + quantization_validation = manifest.get("quantization_validation") + if quantization_validation is not None: + if not isinstance(quantization_validation, dict): + errors.append("quantization_validation manifest entry must be an object") + else: + report_name = str(quantization_validation.get("report", "quantization_report.json")) + report_path = root / report_name + _check_file( + root, + report_name, + checked_files, + errors, + warnings, + required=True, + label="quantization validation report", + ) + if report_path.is_file(): + try: + report = _load_json(report_path) + if report.get("passed") is not True: + errors.append("quantization validation report is not passed") + except (OSError, ValueError, json.JSONDecodeError) as exc: + errors.append(f"invalid quantization validation report: {exc}") + elif family == "rvq" and manifest.get("quantization") == "INT8" and strict_release: + errors.append("missing quantization_validation for INT8 RVQ release") + + if family == "rvq" and manifest.get("quantization") == "INT8": + try: + from compressionkit.export.quantization import require_rvq_encoder_quantization_report + + require_rvq_encoder_quantization_report(root) + except ValueError as exc: + errors.append(str(exc)) + if check_runtime: try: from compressionkit.runtime import load_codec diff --git a/compressionkit/preprocessing/ecg.py b/compressionkit/preprocessing/ecg.py index 3d72fa3..afae0d5 100644 --- a/compressionkit/preprocessing/ecg.py +++ b/compressionkit/preprocessing/ecg.py @@ -17,16 +17,24 @@ from compressionkit.configs.ecg_rvq import AugmentationConfig -def build_preprocessor(frame_size: int, epsilon: float = 1e-3) -> keras.layers.Layer: +def build_preprocessor( + frame_size: int, + epsilon: float = 1e-3, + *, + seed: int | None = None, +) -> keras.layers.Layer: """Create preprocessing pipeline: random crop + layer normalization. Args: frame_size: Number of samples per frame after cropping. epsilon: LayerNorm epsilon for numerical stability. + seed: Optional random-crop seed. Use for reproducible release sampling; + leave unset for training. """ + crop_kwargs = {"seed": seed} if seed is not None else {} return helia.layers.preprocessing.AugmentationPipeline( layers=[ - helia.layers.preprocessing.RandomCrop1D(duration=frame_size, name="RandomCrop"), + helia.layers.preprocessing.RandomCrop1D(duration=frame_size, name="RandomCrop", **crop_kwargs), helia.layers.preprocessing.LayerNormalization1D(epsilon=epsilon, name="LayerNorm"), ] ) diff --git a/compressionkit/preprocessing/ppg.py b/compressionkit/preprocessing/ppg.py index 9190552..3ecc2c0 100644 --- a/compressionkit/preprocessing/ppg.py +++ b/compressionkit/preprocessing/ppg.py @@ -15,16 +15,24 @@ from compressionkit.configs.ppg_rvq import AugmentationConfig -def build_preprocessor(frame_size: int, epsilon: float = 1e-3) -> keras.layers.Layer: +def build_preprocessor( + frame_size: int, + epsilon: float = 1e-3, + *, + seed: int | None = None, +) -> keras.layers.Layer: """Create preprocessing pipeline: random crop + layer normalization. Args: frame_size: Number of samples per frame after cropping. epsilon: LayerNorm epsilon for numerical stability. + seed: Optional random-crop seed. Use for reproducible release sampling; + leave unset for training. """ + crop_kwargs = {"seed": seed} if seed is not None else {} return helia.layers.preprocessing.AugmentationPipeline( layers=[ - helia.layers.preprocessing.RandomCrop1D(duration=frame_size, name="RandomCrop"), + helia.layers.preprocessing.RandomCrop1D(duration=frame_size, name="RandomCrop", **crop_kwargs), helia.layers.preprocessing.LayerNormalization1D(epsilon=epsilon, name="LayerNorm"), ] ) diff --git a/compressionkit/recipes/base_rvq.py b/compressionkit/recipes/base_rvq.py index 009f31a..8951e74 100644 --- a/compressionkit/recipes/base_rvq.py +++ b/compressionkit/recipes/base_rvq.py @@ -28,7 +28,7 @@ from compressionkit.logging.wandb_utils import finalize_wandb_run, init_wandb_run from compressionkit.trainers.common import ( BEST_CKPT_NAME, - collect_rep_dataset, + collect_disjoint_quantization_datasets, extract_history_metrics, reload_best_weights, save_config_snapshot, @@ -270,16 +270,19 @@ def train(self) -> dict[str, Any]: dataset_sources=dataset_sources, ) - rep_dataset = collect_rep_dataset( + rep_dataset, quantization_validation_dataset = collect_disjoint_quantization_datasets( val_ds, - num_batches=self.cfg.evaluation.tflite_rep_batches, - fallback=eval_results["sample_inputs"], + calibration_frames=getattr(self.cfg.evaluation, "int8_calibration_frames", 4096), + validation_frames=getattr(self.cfg.evaluation, "int8_validation_frames", 2048), + sampling_pool_frames=getattr(self.cfg.evaluation, "int8_sampling_pool_frames", 65_536), + seed=getattr(self.cfg.data, "shuffle_seed", 42), ) deploy = export_for_deployment( model.encoder, model.decoder, model.vq.get_weights(), rep_dataset=rep_dataset, + quantization_validation_dataset=quantization_validation_dataset, output_dir=run_dir / "deploy", sample_inputs=eval_results["sample_inputs"], sample_targets=eval_results["sample_targets"], diff --git a/compressionkit/trainers/common.py b/compressionkit/trainers/common.py index 370cd62..179ac1d 100644 --- a/compressionkit/trainers/common.py +++ b/compressionkit/trainers/common.py @@ -101,6 +101,68 @@ def collect_rep_dataset( return fallback.astype(np.float32) +def collect_disjoint_quantization_datasets( + val_ds: tf.data.Dataset, + *, + calibration_frames: int, + validation_frames: int, + sampling_pool_frames: int, + seed: int, +) -> tuple[np.ndarray, np.ndarray]: + """Reservoir-sample disjoint calibration and holdout frame partitions. + + The complete validation split is scanned up to ``sampling_pool_frames``. + Reservoir sampling makes the retained frames representative without holding + the full split in memory, then a seeded shuffle creates non-overlapping + calibration and post-export validation partitions. + + Args: + val_ds: Batched validation dataset yielding ``(inputs, targets)``. + calibration_frames: Number of frames to provide to the TFLite converter. + validation_frames: Number of different frames for post-export parity. + sampling_pool_frames: Maximum source frames to scan before sampling. + seed: Deterministic sampling seed. + + Returns: + ``(calibration_frames, validation_frames)`` arrays. + + Raises: + ValueError: If the validation stream cannot supply both partitions. + """ + required = calibration_frames + validation_frames + if calibration_frames <= 0 or validation_frames <= 0: + raise ValueError("Calibration and validation frame counts must be positive") + if sampling_pool_frames < required: + raise ValueError("sampling_pool_frames must cover both quantization partitions") + + rng = np.random.default_rng(seed) + reservoir: np.ndarray | None = None + frames_seen = 0 + for inputs, _ in val_ds: + batch = np.asarray(inputs.numpy(), dtype=np.float32) + if reservoir is None: + reservoir = np.empty((required, *batch.shape[1:]), dtype=np.float32) + for frame in batch: + if frames_seen < required: + reservoir[frames_seen] = frame + else: + selected = int(rng.integers(0, frames_seen + 1)) + if selected < required: + reservoir[selected] = frame + frames_seen += 1 + if frames_seen >= sampling_pool_frames: + break + if frames_seen >= sampling_pool_frames: + break + + if reservoir is None or frames_seen < required: + raise ValueError( + f"Validation stream yielded {frames_seen} frames; need {required} disjoint quantization frames" + ) + rng.shuffle(reservoir) + return reservoir[:calibration_frames].copy(), reservoir[calibration_frames:].copy() + + # --------------------------------------------------------------------------- # Training history → scalar summary # --------------------------------------------------------------------------- @@ -166,6 +228,7 @@ def write_long_recording_eval(payload: dict[str, Any], run_dir: Path) -> Path: __all__ = [ "BEST_CKPT_NAME", + "collect_disjoint_quantization_datasets", "collect_rep_dataset", "extract_history_metrics", "reload_best_weights", diff --git a/configs/ppg_rvq_smoketest.yaml b/configs/ppg_rvq_smoketest.yaml index 16076be..d905c3e 100644 --- a/configs/ppg_rvq_smoketest.yaml +++ b/configs/ppg_rvq_smoketest.yaml @@ -56,6 +56,9 @@ evaluation: num_samples: 2 input_bit_depth: 16 tflite_rep_batches: 1 + int8_calibration_frames: 4 + int8_validation_frames: 2 + int8_sampling_pool_frames: 6 output: results_root: results/smoketest diff --git a/docs/deployment.md b/docs/deployment.md index 4379f08..d671b10 100644 --- a/docs/deployment.md +++ b/docs/deployment.md @@ -27,6 +27,7 @@ For most users, the only runtime object you need is `compressionkit.runtime.RVQC | `deploy_manifest.json` | Yes | Declares file names, tensor shapes, quantization mode, and codebook metadata. | Runtime metadata | | `encoder.tflite` | Yes | INT8 LiteRT encoder used to produce continuous latents from input frames. | MCU / edge device | | `encoder.h` | Yes | C header for the quantized encoder blob. | MCU firmware | +| `encoder_float32.tflite` | Yes | Float32 LiteRT encoder for browser and host runtimes without Keras. | Browser / x86 / ARM Linux | | `encoder.keras` | Yes | Reference encoder kept in Keras format. | Server / offline tools | | `decoder.keras` | Yes | Reference decoder kept in Keras format. | Server / offline tools | | `decoder_float32.tflite` | Optional, exported by default | Float32 LiteRT decoder for host-side reconstruction without Keras. | x86 / ARM Linux | @@ -35,10 +36,40 @@ For most users, the only runtime object you need is `compressionkit.runtime.RVQC | `codebook.npz` | Yes | NumPy archive containing RVQ codebook tables. | Python runtime | | `codebook.h` | Yes | C header containing RVQ codebook tables. | MCU firmware | | `sample_data.npz` | Optional | Synthetic or evaluation sample inputs / targets / reconstructions. | Validation / demos | +| `demo_recordings.npz` | Release | Ten real, quality-gated continuous recordings, resampled to the model rate. | Browser demos | +| `demo_recordings_manifest.json` | Release | Dataset attribution, license, source offsets, and quality metrics for the demo recordings. | Browser demos / compliance | | `model_card.json` | Optional | Metadata used when publishing to HuggingFace. | Release tooling | The manifest is the source of truth. The runtime reads it first, then resolves the encoder, optional decoder, and codebook files from the names listed there. +## Encoder Preprocessing and INT8 Parity + +Both RVQ encoders accept one normalized frame: PPG is `(1, 1, 320, 1)` at 64 Hz and ECG is `(1, 1, 512, 1)` at 256 Hz. Resample and frame the raw single-channel recording first, then normalize each frame independently: + +```python +mean = frame.mean() +scale = np.sqrt(np.mean((frame - mean) ** 2) + 1e-3) +normalized = (frame - mean) / scale +``` + +Retain `mean` and `scale` to return a decoded frame to its original units: `raw = decoded * scale + mean`. Do not feed raw ADC or continuous values directly to `encoder_int8.tflite`; its calibration domain is this training-time normalized representation. + +Every INT8 RVQ export reservoir-samples up to 65,536 real preprocessed validation frames, using 4,096 frames for LiteRT calibration and a disjoint 2,048-frame holdout for `quantization_report.json`. The release gate compares INT8 and float32 encoder reconstructions on that holdout: at most 1% input saturation and 10% P90 reconstruction PRD. Worst-frame PRD remains a 15% tail-risk warning in the report, rather than a release blocker, because isolated RVQ decision-boundary crossings can be disproportionate. `compressionkit golden validate-all --strict-release` and the Hugging Face publisher reject INT8 RVQ releases without a passed report. + +## Real Browser-Demo Recordings + +RVQ releases carry ten 30-second real recordings in `demo_recordings.npz`: BIDMC PPG at 64 Hz and MIT-BIH ECG at 256 Hz. The companion manifest retains ODC-By attribution and the quality gate results. Signals are continuous source waveforms after resampling; apply the model's normal framing and normalization before inference. + +Refresh local bundles before a manual Hugging Face update: + +```bash +scripts/devcontainer.sh exec -- uv run python scripts/attach_rvq_demo_recordings.py --modality all +``` + +To update only validated release candidates, repeat `--experiment-id`, for example `--experiment-id ecg-rvq-2x --experiment-id ecg-rvq-4x`. + +Use `--duration-seconds 120` for two-minute clips. The command updates only existing `results/*/deploy/` packages, their manifests, and checksums; it does not train or re-export models. + ## Local Runtime Quickstart The local runtime requires only `numpy` and one LiteRT-compatible interpreter package such as `ai-edge-litert`, `tflite-runtime`, or TensorFlow Lite. diff --git a/docs/huggingface.md b/docs/huggingface.md index c51d646..3c86f3a 100644 --- a/docs/huggingface.md +++ b/docs/huggingface.md @@ -23,8 +23,9 @@ for the `snapshot_download` / `from_pretrained` calls below. ## 2. Single-stage codec -Single-stage repos contain `encoder_int8.tflite`, `decoder_int8.tflite`, `codebook.npz`, and -`sample_stimulus.npz`. [`RVQCodec.from_pretrained`](api/models.md) downloads the bundle and wires +Single-stage repos contain `encoder_int8.tflite`, `encoder_float32.tflite`, `decoder_int8.tflite`, +`codebook.npz`, `sample_stimulus.npz`, and `demo_recordings.npz`. `encoder_float32.tflite` is intended for browser and +host LiteRT integrations that need float32 I/O. [`RVQCodec.from_pretrained`](api/models.md) downloads the bundle and wires up the LiteRT interpreters — it needs only NumPy and a LiteRT runtime. ```python @@ -44,6 +45,24 @@ recon = codec.decode(indices) print("shape:", recon.shape) ``` +`demo_recordings.npz` holds ten real, quality-gated continuous examples at the model rate (`signals` has shape `(10, samples)`). Read `demo_recordings_manifest.json` with it: the manifest records source provenance, ODC-By attribution, signal offsets, and the quality measurements used for selection. Apply the model's usual framing and normalization before inference. + +For each 64 Hz PPG 320-sample frame or 256 Hz ECG 512-sample frame, use per-frame layer normalization before either encoder variant: + +```python +mean = frame.mean() +scale = np.sqrt(np.mean((frame - mean) ** 2) + 1e-3) +encoder_input = (frame - mean) / scale +``` + +For display in raw units, undo it after decoding with `decoded * scale + mean`. INT8 RVQ releases produced under the current release policy include `quantization_report.json`, which records parity against `encoder_float32.tflite` on a 2,048-frame real-preprocessed holdout distinct from the 4,096 frames used for LiteRT calibration. + +To refresh all local RVQ bundles before publishing an asset update: + +```bash +scripts/devcontainer.sh exec -- uv run python scripts/attach_rvq_demo_recordings.py --modality all +``` + !!! note Use `RVQCodec.from_pretrained(repo_id)` rather than `RVQCodec(local_dir)` on a raw `snapshot_download` directory. HuggingFace bundles store the manifest as `config.json` diff --git a/docs/methods/rvq.md b/docs/methods/rvq.md index 24eb033..55d3bfd 100644 --- a/docs/methods/rvq.md +++ b/docs/methods/rvq.md @@ -105,5 +105,6 @@ The trained encoder is exported as: - **`encoder.tflite`** — INT8 quantized TFLite model for on-device inference - **`encoder.h`** — C header with the model weights as a byte array +- **`encoder_float32.tflite`** — FP32 LiteRT encoder for browser and host integrations The decoder and RVQ codebooks are stored separately for server-side reconstruction. On-device, only the encoder runs — it produces codebook indices that are transmitted efficiently. diff --git a/docs/release-contract.md b/docs/release-contract.md index 01242af..9004088 100644 --- a/docs/release-contract.md +++ b/docs/release-contract.md @@ -97,6 +97,7 @@ AI packages may expose their demo frames as `sample_data.npz` when the file carr |------|---------| | `encoder.tflite` | Edge encoder | | `encoder.h` | Embedded encoder header | +| `encoder_float32.tflite` | Browser and host LiteRT encoder | | `encoder.keras` | Host reference encode | | `decoder.keras` or `decoder_float32.tflite` | Host reference decode | | `decoder.tflite` and `decoder.h` | Optional on-device decode | diff --git a/scripts/attach_rvq_demo_recordings.py b/scripts/attach_rvq_demo_recordings.py new file mode 100644 index 0000000..bb13bcd --- /dev/null +++ b/scripts/attach_rvq_demo_recordings.py @@ -0,0 +1,141 @@ +"""Attach real, quality-gated ECG and PPG recordings to RVQ golden deploys. + +This is the manual release step for browser demo inputs. It creates the same +10-recording bundle for every compression rate of one modality, then updates +each package manifest and checksum list. It does not train or re-export models. + +Examples: + uv run python scripts/attach_rvq_demo_recordings.py --modality ppg + uv run python scripts/attach_rvq_demo_recordings.py --modality ecg + uv run python scripts/attach_rvq_demo_recordings.py --modality all --duration-seconds 120 +""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +from compressionkit.configs.paths import default_datasets_dir +from compressionkit.experiments.registry import list_goldens +from compressionkit.export.demo_ecg import generate_ecg_demo_clips +from compressionkit.export.demo_ppg import generate_ppg_demo_clips +from compressionkit.export.demo_recordings import attach_demo_recordings_to_deploy, export_demo_recordings +from compressionkit.export.validate import validate_deploy_package + +_SOURCES = { + "ecg": { + "dataset": "MIT-BIH Arrhythmia Database v1.0.0", + "license": "ODC-By-1.0", + "license_url": "https://opendatacommons.org/licenses/by/1-0/", + "url": "https://physionet.org/content/mitdb/1.0.0/", + "citation": "Moody GB, Mark RG. The impact of the MIT-BIH Arrhythmia Database. IEEE Eng Med Biol. 2001;20(3):45-50.", + }, + "ppg": { + "dataset": "BIDMC PPG and Respiration Dataset v1.0.0", + "license": "ODC-By-1.0", + "license_url": "https://opendatacommons.org/licenses/by/1-0/", + "url": "https://physionet.org/content/bidmc/1.0.0/", + "citation": "Pimentel MAF, et al. Toward a robust estimation of respiratory rate from pulse oximeters. IEEE Trans Biomed Eng. 2017;64(8):1914-1923.", + }, +} + + +def _build_clips(modality: str, dataset_root: Path, *, num_clips: int, duration_seconds: float, seed: int): + """Generate the modality's shared demo selection at its golden sample rate.""" + if modality == "ecg": + return generate_ecg_demo_clips( + dataset_root / "mitdb", + num_clips=num_clips, + duration_seconds=duration_seconds, + sample_rate=256, + seed=seed, + ) + return generate_ppg_demo_clips( + dataset_root / "bidmc", + num_clips=num_clips, + duration_seconds=duration_seconds, + sample_rate=64, + seed=seed, + ) + + +def attach_modality( + modality: str, + *, + results_root: Path, + dataset_root: Path, + num_clips: int, + duration_seconds: float, + seed: int, + experiment_ids: set[str] | None = None, +) -> list[Path]: + """Attach one shared real-recordings bundle to selected local RVQ deploys.""" + sample_rate = 256 if modality == "ecg" else 64 + clips = _build_clips( + modality, + dataset_root, + num_clips=num_clips, + duration_seconds=duration_seconds, + seed=seed, + ) + updated: list[Path] = [] + for golden in list_goldens(modality=modality, method="rvq"): + if golden.structure != "codec": + continue + if experiment_ids is not None and golden.experiment_id not in experiment_ids: + continue + deploy_dir = results_root / golden.run_name / "deploy" + if not deploy_dir.is_dir(): + raise FileNotFoundError(f"Missing local deploy package for {golden.experiment_id}: {deploy_dir}") + bundle = export_demo_recordings( + deploy_dir, + clips=clips, + modality=modality, + sample_rate=sample_rate, + source=_SOURCES[modality], + seed=seed, + ) + attach_demo_recordings_to_deploy(deploy_dir, bundle) + result = validate_deploy_package(deploy_dir, strict_release=True) + if not result.ok: + raise RuntimeError(f"Updated {deploy_dir} failed validation: {'; '.join(result.errors)}") + updated.append(deploy_dir) + return updated + + +def main() -> None: + """Parse arguments, attach bundles, and print publish-ready deploy paths.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--modality", choices=["ppg", "ecg", "all"], default="all") + parser.add_argument( + "--experiment-id", + action="append", + default=None, + help="Optional golden experiment ID to update; repeat to update a release subset.", + ) + parser.add_argument("--results-root", type=Path, default=Path("results")) + parser.add_argument("--datasets-root", type=Path, default=Path(default_datasets_dir())) + parser.add_argument("--num-clips", type=int, default=10) + parser.add_argument("--duration-seconds", type=float, default=30.0) + parser.add_argument("--seed", type=int, default=42) + args = parser.parse_args() + + modalities = ("ppg", "ecg") if args.modality == "all" else (args.modality,) + experiment_ids = set(args.experiment_id) if args.experiment_id else None + for modality in modalities: + updated = attach_modality( + modality, + results_root=args.results_root, + dataset_root=args.datasets_root, + num_clips=args.num_clips, + duration_seconds=args.duration_seconds, + seed=args.seed, + experiment_ids=experiment_ids, + ) + print(f"Updated {len(updated)} {modality.upper()} RVQ deploy packages:") + for deploy_dir in updated: + print(f" {deploy_dir}") + + +if __name__ == "__main__": + main() diff --git a/scripts/compare_rvq_encoder_precisions.py b/scripts/compare_rvq_encoder_precisions.py new file mode 100644 index 0000000..904ee47 --- /dev/null +++ b/scripts/compare_rvq_encoder_precisions.py @@ -0,0 +1,166 @@ +"""Compare RVQ encoder precision variants on disjoint real validation data. + +This research script never alters a golden deploy directory. For every selected +RVQ golden, it exports temporary INT8, FP16, and INT16x8 encoders using the +same 4,096-frame calibration partition, then compares each with the float32 +encoder through the shared FP32 RVQ codebook and decoder on a separate +2,048-frame holdout. + +Example: + scripts/devcontainer.sh exec -- uv run python scripts/compare_rvq_encoder_precisions.py +""" + +from __future__ import annotations + +import argparse +import json +import shutil +import tempfile +from dataclasses import asdict +from pathlib import Path + +import keras + +from compressionkit.experiments.registry import GoldenExperiment, list_goldens +from compressionkit.experiments.repackage import collect_release_quantization_frames +from compressionkit.export.quantization import evaluate_rvq_encoder_quantization +from compressionkit.export.tflite import export_encoder_tflite +from compressionkit.runtime._litert import Interpreter + +_PRECISIONS = ("INT8", "FP16", "INT16X8") + + +def _experiment_output( + experiment: GoldenExperiment, + *, + results_root: Path, + precisions: tuple[str, ...], +) -> dict[str, object]: + """Export and score selected encoder precisions for one golden experiment.""" + run_dir = results_root / experiment.run_name + deploy_dir = run_dir / "deploy" + if not deploy_dir.is_dir(): + raise FileNotFoundError(f"Missing deploy directory: {deploy_dir}") + calibration, holdout, contract = collect_release_quantization_frames(experiment, run_dir) + encoder = keras.models.load_model(run_dir / "encoder.keras") + + variants: dict[str, object] = {} + with tempfile.TemporaryDirectory(prefix=f"{experiment.experiment_id}_precision_") as temp_root: + temp_root_path = Path(temp_root) + for precision in precisions: + candidate_dir = temp_root_path / precision.lower() + shutil.copytree(deploy_dir, candidate_dir) + encoder_path, _ = export_encoder_tflite( + encoder, + rep_dataset=calibration, + output_dir=candidate_dir, + tflite_name="encoder.tflite", + header_name="encoder.h", + c_array_name="encoder", + quantization=precision, + io_type="float32" if precision != "INT8" else "int8", + ) + interpreter = Interpreter(model_path=str(encoder_path)) + interpreter.allocate_tensors() + report = evaluate_rvq_encoder_quantization( + candidate_dir, + holdout, + max_frames=None, + ) + variants[precision] = { + "encoder_bytes": encoder_path.stat().st_size, + "input_dtype": str(interpreter.get_input_details()[0]["dtype"]), + "output_dtype": str(interpreter.get_output_details()[0]["dtype"]), + "report": asdict(report), + } + + return { + "modality": experiment.modality, + "sample_rate_hz": experiment.sample_rate, + "compression_ratio": experiment.compression_ratio, + "preprocessing_contract": contract, + "variants": variants, + } + + +def main() -> None: + """Run precision exports and write the comparison JSON.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--modality", choices=("all", "ppg", "ecg"), default="all") + parser.add_argument( + "--experiment", + action="append", + help="Golden experiment ID to evaluate; repeat to select several.", + ) + parser.add_argument("--results-root", type=Path, default=Path("results")) + parser.add_argument( + "--output", + type=Path, + default=Path("results/rvq_encoder_precision_comparison.json"), + help="JSON path for the comparison results.", + ) + parser.add_argument( + "--precision", + action="append", + choices=_PRECISIONS, + help="Precision to test; repeat to select a subset (default: all).", + ) + parser.add_argument( + "--merge", + action="store_true", + help="Merge selected variants into an existing output report.", + ) + args = parser.parse_args() + + precisions = tuple(args.precision or _PRECISIONS) + goldens = [ + golden + for golden in list_goldens(method="rvq") + if golden.structure == "codec" + and (args.modality == "all" or golden.modality == args.modality) + and (not args.experiment or golden.experiment_id in args.experiment) + ] + if args.experiment: + unknown_ids = set(args.experiment) - {golden.experiment_id for golden in goldens} + if unknown_ids: + parser.error(f"Unknown or non-RVQ golden IDs: {', '.join(sorted(unknown_ids))}") + results: dict[str, object] = { + "format_version": 1, + "precisions": list(precisions), + "experiments": {}, + } + if args.merge and args.output.is_file(): + existing = json.loads(args.output.read_text()) + existing_experiments = existing.get("experiments") + if not isinstance(existing_experiments, dict): + parser.error(f"Existing report has invalid experiments field: {args.output}") + existing_precisions = existing.get("precisions", []) + if not isinstance(existing_precisions, list): + parser.error(f"Existing report has invalid precisions field: {args.output}") + results = existing + results["precisions"] = list(dict.fromkeys([*existing_precisions, *precisions])) + for golden in goldens: + print(f"Evaluating {golden.experiment_id}...", flush=True) + current = _experiment_output( + golden, + results_root=args.results_root, + precisions=precisions, + ) + if args.merge and golden.experiment_id in results["experiments"]: + existing = results["experiments"][golden.experiment_id] + if not isinstance(existing, dict): + parser.error(f"Existing report has invalid experiment entry: {golden.experiment_id}") + existing_variants = existing.get("variants") + current_variants = current["variants"] + if not isinstance(existing_variants, dict) or not isinstance(current_variants, dict): + parser.error(f"Existing report has invalid variants: {golden.experiment_id}") + current["variants"] = {**existing_variants, **current_variants} + results["experiments"][golden.experiment_id] = current + + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(json.dumps(results, indent=2, sort_keys=True)) + print(f"Wrote {args.output}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/devcontainer.sh b/scripts/devcontainer.sh index 7693737..4ebc928 100755 --- a/scripts/devcontainer.sh +++ b/scripts/devcontainer.sh @@ -19,6 +19,7 @@ # # Usage: # scripts/devcontainer.sh up # bring the container up (idempotent) +# scripts/devcontainer.sh rebuild # recreate it using the current devcontainer.json # scripts/devcontainer.sh id # print the running container id # scripts/devcontainer.sh exec -- [args] # run a command in the container # scripts/devcontainer.sh shell # open an interactive shell @@ -46,6 +47,16 @@ set -euo pipefail +# This helper manages Docker from the host. A Remote Containers terminal is +# already *in* the target container, where commands should be run directly. +# Detect that case before attempting to discover the host's NVM installation; +# otherwise the resulting "devcontainer CLI not found" error is misleading. +if [[ -f /.dockerenv || -f /run/.containerenv ]]; then + echo "error: scripts/devcontainer.sh must be run from the host, not from inside the dev container." >&2 + echo "Run the requested command directly here instead (for example: uv run pytest -q)." >&2 + exit 2 +fi + # Resolve the main worktree path: the dev container is created against the # primary checkout, not any linked `git worktree add` copy. `git worktree # list` always prints the main worktree first. @@ -65,6 +76,52 @@ if [[ -z "${MAIN_WORKSPACE}" ]]; then exit 1 fi +# Outside-agent shells are non-interactive, so they do not necessarily source +# the user's shell profile that activates an NVM-managed Node installation. +# The Dev Containers CLI is a Node executable whose ``#!/usr/bin/env node`` +# shebang must use the same Node version that installed it. If it is not +# already on PATH, find the newest NVM-managed CLI and prepend its ``bin`` +# directory, which makes both ``devcontainer`` and its matching ``node`` +# available without changing the caller's global Node selection. +ensure_devcontainer_cli() { + if command -v devcontainer >/dev/null 2>&1; then + return 0 + fi + + local nvm_dir="${NVM_DIR:-${HOME:-}/.nvm}" + local candidates=() + local candidate + shopt -s nullglob + candidates=("${nvm_dir}"/versions/node/*/bin/devcontainer) + shopt -u nullglob + + for ((candidate = ${#candidates[@]} - 1; candidate >= 0; candidate--)); do + if [[ -x "${candidates[candidate]}" ]]; then + export PATH="$(dirname "${candidates[candidate]}"):${PATH}" + return 0 + fi + done + + echo "error: devcontainer CLI not found on PATH or under ${nvm_dir}" >&2 + echo "Install @devcontainers/cli or set PATH/NVM_DIR before running this helper." >&2 + return 1 +} + +ensure_devcontainer_cli + +acquire_lifecycle_lock() { + # Two overlapping `devcontainer up` invocations can both observe no + # container and create duplicates with the same workspace label. Keep + # the lock host-local because this script controls local Docker state. + local workspace_hash + workspace_hash="$(printf '%s' "${MAIN_WORKSPACE}" | sha256sum | cut -c1-16)" + exec 9>"/tmp/devcontainer-${workspace_hash}.lock" + if ! flock -n 9; then + echo "error: another dev-container up or rebuild is already running for ${MAIN_WORKSPACE}" >&2 + return 1 + fi +} + container_id() { local ids ids="$(docker ps -q --filter "label=devcontainer.local_folder=${MAIN_WORKSPACE}")" @@ -78,6 +135,18 @@ container_id() { } cmd_up() { + acquire_lifecycle_lock + devcontainer up --workspace-folder "${MAIN_WORKSPACE}" +} + +cmd_rebuild() { + acquire_lifecycle_lock + local id + id="$(container_id)" + if [[ -n "${id}" ]]; then + echo "Removing the existing dev container so the current configuration is applied..." + docker rm -f "${id}" >/dev/null + fi devcontainer up --workspace-folder "${MAIN_WORKSPACE}" } @@ -187,6 +256,7 @@ subcommand="${1:-}" case "${subcommand}" in up) cmd_up ;; + rebuild) cmd_rebuild ;; id) cmd_id ;; exec) [[ "${1:-}" == "--" ]] && shift @@ -198,7 +268,7 @@ case "${subcommand}" in gpu-check) cmd_gpu_check ;; gpu-recover) cmd_gpu_recover ;; *) - echo "usage: $(basename "$0") {up|id|exec -- |shell|down|status|gpu-check|gpu-recover}" >&2 + echo "usage: $(basename "$0") {up|rebuild|id|exec -- |shell|down|status|gpu-check|gpu-recover}" >&2 exit 1 ;; esac diff --git a/scripts/generate_ecg_demo_samples.py b/scripts/generate_ecg_demo_samples.py new file mode 100644 index 0000000..a26bce9 --- /dev/null +++ b/scripts/generate_ecg_demo_samples.py @@ -0,0 +1,51 @@ +"""Generate quality-gated, real MIT-BIH ECG CSV files for the web demo. + +Examples: + uv run python scripts/generate_ecg_demo_samples.py + uv run python scripts/generate_ecg_demo_samples.py --duration-seconds 120 +""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +from compressionkit.configs.paths import default_datasets_dir +from compressionkit.export.demo_ecg import export_ecg_demo_csvs + + +def main() -> None: + """Parse arguments and write ECG demo CSV artifacts.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--output-dir", + type=Path, + default=Path("results/demo_ecg_real_mitdb_256hz_30s"), + help="Destination directory for CSV files and manifest.json.", + ) + parser.add_argument( + "--dataset-dir", + type=Path, + default=Path(default_datasets_dir()) / "mitdb", + help="Canonical MIT-BIH H5 directory.", + ) + parser.add_argument("--num-clips", type=int, default=10, help="Number of clips to generate (1-10).") + parser.add_argument("--duration-seconds", type=float, default=30.0, help="Length of each clip (minimum 10).") + parser.add_argument("--sample-rate", type=int, default=256, help="Output sample rate in Hz.") + parser.add_argument("--seed", type=int, default=42, help="Reproducibility seed.") + args = parser.parse_args() + + exported = export_ecg_demo_csvs( + args.output_dir, + dataset_dir=args.dataset_dir, + num_clips=args.num_clips, + duration_seconds=args.duration_seconds, + sample_rate=args.sample_rate, + seed=args.seed, + ) + print(f"Wrote {len(exported.csv_paths)} ECG CSV files to {args.output_dir}") + print(f"Quality manifest: {exported.manifest_path}") + + +if __name__ == "__main__": + main() diff --git a/scripts/publish_to_huggingface.py b/scripts/publish_to_huggingface.py index a3112e8..753794f 100644 --- a/scripts/publish_to_huggingface.py +++ b/scripts/publish_to_huggingface.py @@ -34,6 +34,7 @@ from compressionkit.export.artifact_contract import ArtifactFile from compressionkit.export.family_registry import CodecFamilySpec, get_family_spec +from compressionkit.export.quantization import require_rvq_encoder_quantization_report logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s") logger = logging.getLogger(__name__) @@ -132,6 +133,7 @@ def publish( family = _detect_family(deploy_dir) spec = get_family_spec(family) + require_rvq_encoder_quantization_report(deploy_dir) logger.info("Detected deploy family: %s", family) # Validate HuggingFace availability before allocating any resources (skip for dry runs) diff --git a/tests/test_base_rvq_scorecard_sync.py b/tests/test_base_rvq_scorecard_sync.py index d985957..18c11bb 100644 --- a/tests/test_base_rvq_scorecard_sync.py +++ b/tests/test_base_rvq_scorecard_sync.py @@ -71,7 +71,11 @@ def test_base_rvq_trainer_syncs_built_scorecard_into_deploy(monkeypatch, tmp_pat model_dump=lambda: {"run_name": "ppg_rvq_64hz_04x_golden"}, data=SimpleNamespace(sampling_rate=64, steps_per_epoch=1, epochs=1), model=SimpleNamespace(kmeans_init=False), - evaluation=SimpleNamespace(tflite_rep_batches=1), + evaluation=SimpleNamespace( + int8_calibration_frames=1, + int8_validation_frames=1, + int8_sampling_pool_frames=2, + ), training=SimpleNamespace(selection_metric="val_loss"), output=SimpleNamespace( results_root=tmp_path, log_file="train.log", wandb=SimpleNamespace(artifact_summary_only=False) @@ -90,7 +94,11 @@ def test_base_rvq_trainer_syncs_built_scorecard_into_deploy(monkeypatch, tmp_pat monkeypatch.setattr(base_rvq, "build_callbacks", lambda *args, **kwargs: []) monkeypatch.setattr(base_rvq, "reload_best_weights", lambda *args, **kwargs: None) monkeypatch.setattr(base_rvq, "save_model_artifacts", lambda *args, **kwargs: {}) - monkeypatch.setattr(base_rvq, "collect_rep_dataset", lambda *args, **kwargs: np.zeros((1, 4, 1), dtype=np.float32)) + monkeypatch.setattr( + base_rvq, + "collect_disjoint_quantization_datasets", + lambda *args, **kwargs: (np.zeros((1, 4, 1), dtype=np.float32), np.ones((1, 4, 1), dtype=np.float32)), + ) monkeypatch.setattr( base_rvq, "export_for_deployment", diff --git a/tests/test_demo_ecg.py b/tests/test_demo_ecg.py new file mode 100644 index 0000000..3ab3c6b --- /dev/null +++ b/tests/test_demo_ecg.py @@ -0,0 +1,72 @@ +"""Tests for browser-demo ECG sample generation.""" + +from __future__ import annotations + +import json + +import h5py +import numpy as np + +from compressionkit.export.demo_ecg import export_ecg_demo_csvs, generate_ecg_demo_clips +from compressionkit.preprocessing.ecg import generate_synthetic_ecg_batch + + +def _write_mitdb_fixture(root, *, count: int = 2) -> None: + """Write canonical MIT-BIH-shaped H5 records for source-selection tests.""" + for index in range(count): + signal = generate_synthetic_ecg_batch( + num_segments=1, + signal_length=360 * 20, + sample_rate=360, + noise_multiplier=[0.05, 0.05], + seed=index + 1, + )[0] + with h5py.File(root / f"{100 + index}.h5", "w") as handle: + handle.create_dataset("data", data=np.stack([signal, signal])) + handle.attrs["fs"] = 360 + handle.attrs["source"] = "mitdb" + handle.attrs["patient_id"] = str(100 + index) + + +def test_generate_ecg_demo_clips_are_quality_gated(tmp_path) -> None: + """Generated demo clips are continuous 256 Hz ECG signals with valid rhythm metrics.""" + _write_mitdb_fixture(tmp_path) + clips = generate_ecg_demo_clips(tmp_path, num_clips=2, duration_seconds=10, sample_rate=256, seed=7) + + assert len(clips) == 2 + for clip in clips: + assert clip.signal.shape == (2560,) + assert clip.signal.dtype == np.float32 + assert 45.0 <= clip.quality.heart_rate_bpm <= 110.0 + assert clip.quality.num_r_peaks >= 8 + assert clip.quality.rr_cv <= 0.15 + assert clip.quality.clipping_fraction < 0.01 + assert clip.source_record in {"100", "101"} + assert clip.source_sample_rate == 360 + + +def test_export_ecg_demo_csvs_writes_manifest(tmp_path) -> None: + """CSV export writes an inspectable clip manifest with quality metadata.""" + source_dir = tmp_path / "mitdb" + source_dir.mkdir() + _write_mitdb_fixture(source_dir) + exported = export_ecg_demo_csvs( + tmp_path / "output", + dataset_dir=source_dir, + num_clips=2, + duration_seconds=10, + sample_rate=256, + seed=7, + ) + + assert len(exported.csv_paths) == 2 + rows = np.loadtxt(exported.csv_paths[0], delimiter=",", skiprows=1) + assert rows.shape == (2560, 3) + assert np.allclose(rows[:, 1], np.arange(2560) / 256.0) + + manifest = json.loads(exported.manifest_path.read_text()) + assert manifest["modality"] == "ecg" + assert manifest["source"]["license"] == "ODC-By-1.0" + assert manifest["sample_rate"] == 256 + assert len(manifest["clips"]) == 2 + assert manifest["clips"][0]["quality"]["num_r_peaks"] >= 8 diff --git a/tests/test_demo_ppg.py b/tests/test_demo_ppg.py new file mode 100644 index 0000000..c7cbf8f --- /dev/null +++ b/tests/test_demo_ppg.py @@ -0,0 +1,39 @@ +"""Tests for real BIDMC PPG browser-demo selection.""" + +from __future__ import annotations + +import h5py +import numpy as np + +from compressionkit.export.demo_ppg import generate_ppg_demo_clips + + +def _write_bidmc_fixture(root, *, count: int = 2) -> None: + """Write canonical BIDMC-shaped H5 records with clean pulse waveforms.""" + sample_rate = 125 + times = np.arange(sample_rate * 20, dtype=np.float32) / sample_rate + for index in range(count): + frequency = 1.2 + index * 0.1 + signal = np.sin(2 * np.pi * frequency * times) + 0.3 * np.sin(4 * np.pi * frequency * times) + with h5py.File(root / f"bidmc{index:02d}.h5", "w") as handle: + handle.create_dataset("data", data=signal[np.newaxis, :]) + handle.attrs["fs"] = sample_rate + handle.attrs["source"] = "bidmc" + handle.attrs["patient_id"] = f"bidmc{index:02d}" + + +def test_generate_ppg_demo_clips_are_quality_gated(tmp_path) -> None: + """Selected PPG clips are continuous, resampled, and physiologically plausible.""" + _write_bidmc_fixture(tmp_path) + + clips = generate_ppg_demo_clips(tmp_path, num_clips=2, duration_seconds=10, sample_rate=64, seed=7) + + assert len(clips) == 2 + for clip in clips: + assert clip.signal.shape == (640,) + assert clip.signal.dtype == np.float32 + assert 45.0 <= clip.quality.heart_rate_bpm <= 120.0 + assert clip.quality.heart_rate_qos >= 0.8 + assert clip.quality.pulse_snr_db >= 8.0 + assert clip.quality.rr_cv <= 0.2 + assert clip.source_sample_rate == 125 diff --git a/tests/test_demo_recordings.py b/tests/test_demo_recordings.py new file mode 100644 index 0000000..922d9ce --- /dev/null +++ b/tests/test_demo_recordings.py @@ -0,0 +1,57 @@ +"""Tests for deployable real-recordings demo bundles.""" + +from __future__ import annotations + +import json +from dataclasses import dataclass + +import numpy as np + +from compressionkit.export.artifact_contract import DemoArray +from compressionkit.export.demo_recordings import attach_demo_recordings_to_deploy, export_demo_recordings + + +@dataclass(frozen=True) +class _Quality: + score: float + + +@dataclass(frozen=True) +class _Clip: + signal: np.ndarray + quality: _Quality + source_record: str + source_sample_rate: int + start_seconds: float + + +def test_export_and_attach_demo_recordings(tmp_path) -> None: + """The public bundle has stable arrays, attribution, manifest registration, and checksums.""" + deploy_dir = tmp_path / "deploy" + deploy_dir.mkdir() + (deploy_dir / "deploy_manifest.json").write_text(json.dumps({"artifacts": {}})) + clips = [ + _Clip(np.arange(20, dtype=np.float32), _Quality(score=0.9), "subject-a", 125, 3.5), + _Clip(np.arange(20, dtype=np.float32) + 1, _Quality(score=0.95), "subject-b", 125, 6.5), + ] + + bundle = export_demo_recordings( + deploy_dir, + clips=clips, + modality="ppg", + sample_rate=64, + source={"dataset": "BIDMC", "license": "ODC-By-1.0"}, + seed=42, + ) + attach_demo_recordings_to_deploy(deploy_dir, bundle) + + with np.load(bundle.recordings_path) as recordings: + assert recordings[DemoArray.SIGNALS].shape == (2, 20) + assert int(recordings[DemoArray.SAMPLE_RATE]) == 64 + assert recordings[DemoArray.SOURCE_RECORDS].tolist() == ["subject-a", "subject-b"] + manifest = json.loads(bundle.manifest_path.read_text()) + assert manifest["source"]["license"] == "ODC-By-1.0" + assert manifest["recordings"][0]["quality"]["score"] == 0.9 + deploy_manifest = json.loads((deploy_dir / "deploy_manifest.json").read_text()) + assert deploy_manifest["artifacts"]["demo_recordings"] == "demo_recordings.npz" + assert (deploy_dir / "checksums.json").is_file() diff --git a/tests/test_deploy_enhanced.py b/tests/test_deploy_enhanced.py index b6f653f..652b866 100644 --- a/tests/test_deploy_enhanced.py +++ b/tests/test_deploy_enhanced.py @@ -65,6 +65,9 @@ def test_new_fields_exist(self): from compressionkit.export.deploy import DeploymentArtifacts arts = DeploymentArtifacts(output_dir=Path("/tmp/test")) + assert hasattr(arts, "encoder_float32_tflite") + assert hasattr(arts, "encoder_fp16_tflite") + assert hasattr(arts, "encoder_int16x8_tflite") assert hasattr(arts, "decoder_float32_tflite") assert hasattr(arts, "decoder_int8_tflite") assert hasattr(arts, "decoder_int8_header") @@ -77,6 +80,9 @@ def test_as_dict_includes_new_fields(self): arts = DeploymentArtifacts(output_dir=Path("/tmp/test")) d = arts.as_dict() + assert "encoder_float32_tflite" in d + assert "encoder_fp16_tflite" in d + assert "encoder_int16x8_tflite" in d assert "decoder_float32_tflite" in d assert "model_card" in d assert "scorecard" in d diff --git a/tests/test_huggingface.py b/tests/test_huggingface.py index 82aab92..24cfd51 100644 --- a/tests/test_huggingface.py +++ b/tests/test_huggingface.py @@ -21,6 +21,7 @@ def mock_deploy(tmp_path: Path) -> Path: manifest = { "model_name": "ppg_rvq_64hz_04x_test", + "family": "rvq", "model_version": "1.0", "quantization": "INT8", "io_type": "int8", @@ -39,8 +40,25 @@ def mock_deploy(tmp_path: Path) -> Path: "embedding_dim": 16, }, "sample_data": {"npz": "sample_data.npz", "num_samples": 10, "arrays": ["inputs"]}, + "quantization_validation": {"report": "quantization_report.json", "passed": True}, } (deploy / "deploy_manifest.json").write_text(json.dumps(manifest)) + (deploy / "quantization_report.json").write_text( + json.dumps( + { + "passed": True, + "metrics": { + "code_index_match_fraction_median": 0.5, + "encoder_input_saturation_fraction_max": 0.0, + "encoder_input_saturation_fraction_p90": 0.0, + "frames_checked": 128, + "latent_prd_percent_p90": 1.0, + "reconstruction_prd_percent_max": 1.0, + "reconstruction_prd_percent_p90": 1.0, + }, + } + ) + ) return deploy @@ -164,6 +182,7 @@ class TestPublishStaging: def test_dry_run_stages_files(self, mock_deploy, mock_scorecard): # Create some dummy files in the deploy dir (mock_deploy / "encoder.tflite").write_bytes(b"\x00" * 100) + (mock_deploy / "encoder_float32.tflite").write_bytes(b"\x00" * 100) (mock_deploy / "codebook.npz").write_bytes(b"\x00" * 50) (mock_deploy / "encoder.h").write_text("// header") @@ -180,10 +199,12 @@ def test_dry_run_stages_files(self, mock_deploy, mock_scorecard): # Check renamed files assert (staging_dir / "encoder_int8.tflite").exists() + assert (staging_dir / "encoder_float32.tflite").exists() assert (staging_dir / "codebook.npz").exists() assert (staging_dir / "encoder.h").exists() assert (staging_dir / "config.json").exists() # renamed from deploy_manifest.json assert (staging_dir / "quality_scorecard.json").exists() + assert (staging_dir / "quantization_report.json").exists() assert (staging_dir / "README.md").exists() # README should be a proper model card @@ -200,6 +221,17 @@ def test_dry_run_no_manifest_exits(self, tmp_path): with pytest.raises(FileNotFoundError): publish(deploy_dir=tmp_path, repo_id="test/test", dry_run=True) + def test_dry_run_rejects_int8_rvq_without_quantization_report(self, mock_deploy): + from scripts.publish_to_huggingface import publish + + manifest_path = mock_deploy / "deploy_manifest.json" + manifest = json.loads(manifest_path.read_text()) + manifest.pop("quantization_validation") + manifest_path.write_text(json.dumps(manifest)) + + with pytest.raises(ValueError, match="quantization_validation"): + publish(deploy_dir=mock_deploy, repo_id="test/test", dry_run=True) + def test_dry_run_hybrid_stages_denoiser(self, tmp_path): """A hybrid deploy (DSP backend + learned denoiser) must publish the denoiser. diff --git a/tests/test_rvq_deploy_smoke.py b/tests/test_rvq_deploy_smoke.py index f114bf7..46a9d45 100644 --- a/tests/test_rvq_deploy_smoke.py +++ b/tests/test_rvq_deploy_smoke.py @@ -2,6 +2,8 @@ from __future__ import annotations +import json + import keras import numpy as np @@ -80,6 +82,11 @@ def test_export_for_deployment_smoke_validates_strict_release(tmp_path) -> None: assert artifacts.manifest.exists() assert artifacts.codec_spec.exists() + assert artifacts.encoder_float32_tflite.exists() + assert artifacts.encoder_fp16_tflite.exists() + assert artifacts.encoder_fp16_header.exists() + assert artifacts.encoder_int16x8_tflite.exists() + assert artifacts.encoder_int16x8_header.exists() assert artifacts.scorecard.exists() assert artifacts.reference_vectors.exists() assert artifacts.readme.exists() @@ -95,6 +102,26 @@ def test_export_for_deployment_smoke_validates_strict_release(tmp_path) -> None: assert "scorecard.json" in result.checked_files assert "reference_vectors.npz" in result.checked_files + manifest = json.loads(artifacts.manifest.read_text()) + assert manifest["encoder"]["default_variant"] == "float32" + assert manifest["encoder"]["default_tflite"] == "encoder_float32.tflite" + assert manifest["encoder"]["float32_tflite"] == "encoder_float32.tflite" + assert manifest["encoder"]["fp16_tflite"] == "encoder_fp16.tflite" + assert manifest["encoder"]["int16x8_tflite"] == "encoder_int16x8.tflite" + + from compressionkit.runtime._litert import Interpreter + + interpreter = Interpreter(model_path=str(artifacts.encoder_float32_tflite)) + interpreter.allocate_tensors() + input_detail = interpreter.get_input_details()[0] + output_detail = interpreter.get_output_details()[0] + assert input_detail["dtype"] == np.float32 + assert output_detail["dtype"] == np.float32 + sample = np.load(deploy_dir / "sample_data.npz")["inputs"][:1] + interpreter.set_tensor(input_detail["index"], sample) + interpreter.invoke() + assert interpreter.get_tensor(output_detail["index"]).dtype == np.float32 + def test_validate_deploy_cli_smoke(tmp_path, capsys) -> None: deploy_dir, _artifacts = _build_smoke_rvq_deploy(tmp_path) diff --git a/tests/test_rvq_quantization.py b/tests/test_rvq_quantization.py new file mode 100644 index 0000000..1e818a9 --- /dev/null +++ b/tests/test_rvq_quantization.py @@ -0,0 +1,95 @@ +"""Tests for INT8 RVQ encoder release gates.""" + +from __future__ import annotations + +import json +from pathlib import Path + +import numpy as np +import pytest +import tensorflow as tf + +from compressionkit.export.quantization import ( + RvqEncoderQuantizationReport, + require_rvq_encoder_quantization_report, +) +from compressionkit.trainers.common import collect_disjoint_quantization_datasets + + +def _report(*, p90: float = 8.0, maximum: float = 14.0, saturation: float = 0.0) -> RvqEncoderQuantizationReport: + """Build a representative parity result for threshold tests.""" + return RvqEncoderQuantizationReport( + code_index_match_fraction_median=0.5, + encoder_input_saturation_fraction_max=saturation, + encoder_input_saturation_fraction_p90=0.0, + frames_checked=128, + latent_prd_percent_p90=3.0, + reconstruction_prd_percent_p90=p90, + reconstruction_prd_percent_max=maximum, + ) + + +def test_quantization_report_accepts_reconstruction_parity_without_index_equality() -> None: + """Equivalent RVQ entries need not produce identical index streams.""" + assert _report().passed() + + +@pytest.mark.parametrize( + ("kwargs", "expected"), + [ + ({"saturation": 0.0101}, "saturation"), + ], +) +def test_quantization_report_rejects_saturation(kwargs: dict[str, float], expected: str) -> None: + """The publication gate rejects invalid INT8 input ranges.""" + assert not _report(**kwargs).passed(), expected + + +def test_quantization_report_keeps_worst_frame_as_diagnostic() -> None: + """An isolated RVQ decision-boundary crossing does not block release.""" + assert _report(maximum=80.0).passed() + + +def test_quantization_report_keeps_p90_as_recommendation() -> None: + """High P90 is reported to users rather than suppressing the release.""" + assert _report(p90=80.0).passed() + + +def test_publisher_guard_requires_passed_report_for_int8_rvq(tmp_path: Path) -> None: + """Hub releases cannot bypass the parity gate with a hand-written manifest.""" + (tmp_path / "deploy_manifest.json").write_text( + json.dumps( + { + "family": "rvq", + "quantization": "INT8", + "quantization_validation": {"report": "quantization_report.json", "passed": True}, + } + ) + ) + (tmp_path / "quantization_report.json").write_text( + json.dumps({"passed": True, "metrics": _report().__dict__}) + ) + + require_rvq_encoder_quantization_report(tmp_path) + + (tmp_path / "quantization_report.json").write_text(json.dumps({"passed": False, "metrics": _report().__dict__})) + with pytest.raises(ValueError, match="not passed"): + require_rvq_encoder_quantization_report(tmp_path) + + +def test_quantization_partitions_are_disjoint_and_seeded() -> None: + """Calibration ranges and parity measurements must not reuse a frame.""" + frames = np.arange(20, dtype=np.float32).reshape(20, 1, 1, 1) + dataset = tf.data.Dataset.from_tensor_slices((frames, frames)).batch(4) + + calibration, validation = collect_disjoint_quantization_datasets( + dataset, + calibration_frames=8, + validation_frames=6, + sampling_pool_frames=20, + seed=123, + ) + + assert calibration.shape[0] == 8 + assert validation.shape[0] == 6 + assert set(calibration[:, 0, 0, 0]).isdisjoint(validation[:, 0, 0, 0]) From ea34d6793c4a87de236bc3e514400efbf0f9dda1 Mon Sep 17 00:00:00 2001 From: Adam Page Date: Tue, 21 Jul 2026 19:44:26 -0500 Subject: [PATCH 2/4] fix: preserve encoder precision reports --- compressionkit/export/release.py | 1 + 1 file changed, 1 insertion(+) diff --git a/compressionkit/export/release.py b/compressionkit/export/release.py index 85f2f06..9ec5849 100644 --- a/compressionkit/export/release.py +++ b/compressionkit/export/release.py @@ -83,6 +83,7 @@ def build_model_card( "license": model_card_info.get("license", "other"), "dataset_sources": model_card_info.get("dataset_sources"), "scorecard_summary": model_card_info.get("scorecard_summary", {}), + "encoder_precision_report": model_card_info.get("encoder_precision_report", {}), } From 89f2c12a5339a9a71193882a2be7f39784d65b63 Mon Sep 17 00:00:00 2001 From: Adam Page Date: Tue, 21 Jul 2026 20:36:34 -0500 Subject: [PATCH 3/4] style: format release helpers --- compressionkit/experiments/repackage.py | 4 +--- compressionkit/export/quantization.py | 9 +++++---- tests/test_rvq_quantization.py | 4 +--- 3 files changed, 7 insertions(+), 10 deletions(-) diff --git a/compressionkit/experiments/repackage.py b/compressionkit/experiments/repackage.py index 08d2b29..c1ddfb1 100644 --- a/compressionkit/experiments/repackage.py +++ b/compressionkit/experiments/repackage.py @@ -192,9 +192,7 @@ def repackage_rvq_golden( variants = precision_payload.get("experiments", {}).get(experiment.experiment_id, {}).get("variants", {}) if isinstance(variants, dict): model_card_info["encoder_precision_report"] = { - name: payload.get("report", {}) - for name, payload in variants.items() - if isinstance(payload, dict) + name: payload.get("report", {}) for name, payload in variants.items() if isinstance(payload, dict) } break if scorecard_payload is not None: diff --git a/compressionkit/export/quantization.py b/compressionkit/export/quantization.py index b4a340a..10526ad 100644 --- a/compressionkit/export/quantization.py +++ b/compressionkit/export/quantization.py @@ -37,9 +37,7 @@ def passed(self) -> bool: worst-frame PRD are retained as explicit quality recommendations so customers can select an encoder precision with full visibility. """ - return ( - self.encoder_input_saturation_fraction_max <= MAX_INPUT_SATURATION_FRACTION - ) + return self.encoder_input_saturation_fraction_max <= MAX_INPUT_SATURATION_FRACTION def _prd_percent(actual: np.ndarray, expected: np.ndarray) -> float: @@ -83,7 +81,10 @@ def evaluate_rvq_encoder_quantization( float_encoder = Interpreter(model_path=str(root / ArtifactFile.ENCODER_FLOAT32_TFLITE)) candidate_encoder.allocate_tensors() float_encoder.allocate_tensors() - candidate_input, candidate_output = candidate_encoder.get_input_details()[0], candidate_encoder.get_output_details()[0] + candidate_input, candidate_output = ( + candidate_encoder.get_input_details()[0], + candidate_encoder.get_output_details()[0], + ) float_input, float_output = float_encoder.get_input_details()[0], float_encoder.get_output_details()[0] codec = RVQCodec(root) diff --git a/tests/test_rvq_quantization.py b/tests/test_rvq_quantization.py index 1e818a9..3e5ec8d 100644 --- a/tests/test_rvq_quantization.py +++ b/tests/test_rvq_quantization.py @@ -66,9 +66,7 @@ def test_publisher_guard_requires_passed_report_for_int8_rvq(tmp_path: Path) -> } ) ) - (tmp_path / "quantization_report.json").write_text( - json.dumps({"passed": True, "metrics": _report().__dict__}) - ) + (tmp_path / "quantization_report.json").write_text(json.dumps({"passed": True, "metrics": _report().__dict__})) require_rvq_encoder_quantization_report(tmp_path) From 7fb36636498af6b053e1451a87286a2e699d3874 Mon Sep 17 00:00:00 2001 From: Adam Page Date: Tue, 21 Jul 2026 20:52:05 -0500 Subject: [PATCH 4/4] test: cover FP32 RVQ contract artifact --- tests/test_deploy_contract.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/test_deploy_contract.py b/tests/test_deploy_contract.py index bf5681b..35be1a4 100644 --- a/tests/test_deploy_contract.py +++ b/tests/test_deploy_contract.py @@ -52,6 +52,7 @@ def test_strict_rvq_contract_accepts_complete_file_set(tmp_path) -> None: _write_minimal_manifest(tmp_path, family="rvq") for rel in [ "encoder.tflite", + "encoder_float32.tflite", "encoder.keras", "decoder.keras", "codebook.npz",