Download loader/prism_quant/runtime.py from atomtanstudio/Prism-Q6: direct link, hf CLI and curl.
- Browser
- Download file 43.3 kB
-
https://huggingface.co/atomtanstudio/Prism-Q6/resolve/main/loader/prism_quant/runtime.py
- Command line
-
hf download hf://atomtanstudio/Prism-Q6/loader/prism_quant/runtime.py
-
curl -L -o runtime.py https://huggingface.co/atomtanstudio/Prism-Q6/resolve/main/loader/prism_quant/runtime.py
43.3 kB
| """Isolated, CPU-backed inference around the pinned official Prism pipeline. | |
| Importing this module does not import Prism, CUDA kernels, or model dependencies. | |
| The official forward, sparse attention and scheduler remain upstream code. | |
| """ | |
| from __future__ import annotations | |
| import contextlib | |
| import functools | |
| import hashlib | |
| import json | |
| import math | |
| import os | |
| from pathlib import Path | |
| import resource | |
| import subprocess | |
| import sys | |
| import tempfile | |
| import time | |
| import wave | |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] | |
| def atomic_json(path, value): | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with tempfile.NamedTemporaryFile("w", dir=path.parent, delete=False, suffix=".tmp") as stream: | |
| temporary = Path(stream.name) | |
| try: | |
| json.dump(value, stream, indent=2, allow_nan=False) | |
| stream.write("\n") | |
| except BaseException: | |
| temporary.unlink(missing_ok=True) | |
| raise | |
| temporary.replace(path) | |
| def read_json(path): | |
| with Path(path).open() as stream: | |
| value = json.load(stream) | |
| if not isinstance(value, dict): | |
| raise ValueError(f"Expected JSON object: {path}") | |
| return value | |
| def validate_settings(height, width, frames, fps, steps, spatial_scale=8, temporal_scale=4): | |
| if height <= 0 or width <= 0 or height % (spatial_scale * 2) or width % (spatial_scale * 2): | |
| raise ValueError(f"Height and width must be positive multiples of {spatial_scale * 2}") | |
| if frames < temporal_scale + 1 or (frames - 1) % temporal_scale: | |
| raise ValueError(f"Frames must be >= {temporal_scale + 1} and (frames - 1) divisible by {temporal_scale}") | |
| if not math.isfinite(fps) or fps <= 0 or steps < 2: | |
| raise ValueError("FPS must be finite and positive; use at least two diffusion steps") | |
| def prepared_source(vendor_dir): | |
| """Reject the known remote-kernel bootstrap before importing any upstream code.""" | |
| vendor_dir = Path(vendor_dir).resolve() | |
| source = vendor_dir / "hymm/models/modules/wan_video_dit.py" | |
| if not source.is_file(): | |
| raise FileNotFoundError(f"Prism source is missing: {source}") | |
| if "get_kernel(" in source.read_text(): | |
| raise RuntimeError("Run scripts/prepare_source.py before inference: remote kernel loader remains") | |
| if str(vendor_dir) not in sys.path: | |
| sys.path.insert(0, str(vendor_dir)) | |
| return vendor_dir | |
| def prepared_guidance_source(vendor_dir): | |
| """Verify the independently reviewed guidance patch before pipeline import.""" | |
| from scripts.prepare_guidance_patch import prepare | |
| receipt = prepare(vendor_dir, check_only=True) | |
| if not receipt["already_prepared"]: | |
| raise RuntimeError("Run scripts/prepare_guidance_patch.py before inference") | |
| return receipt | |
| def resolve_cfg(cfg=5.0, visual_cfg=None, audio_cfg=None): | |
| values = (cfg, cfg if visual_cfg is None else visual_cfg, cfg if audio_cfg is None else audio_cfg) | |
| if any(not math.isfinite(value) or value <= 0 for value in values): | |
| raise ValueError("CFG scales must be finite and positive") | |
| return {"fallback": float(values[0]), "video": float(values[1]), "audio": float(values[2])} | |
| class RunTelemetry: | |
| def __init__(self, device=None): | |
| self.device = device | |
| self.events = [] | |
| self.started = time.monotonic() | |
| def memory(self): | |
| # ru_maxrss is KiB on Linux and bytes on macOS. | |
| rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss | |
| result = {"host_peak_rss_bytes": int(rss if sys.platform == "darwin" else rss * 1024)} | |
| if self.device is not None: | |
| import torch | |
| if str(self.device).startswith("cuda") and torch.cuda.is_available(): | |
| result.update( | |
| cuda_allocated_bytes=torch.cuda.memory_allocated(self.device), | |
| cuda_reserved_bytes=torch.cuda.memory_reserved(self.device), | |
| cuda_peak_allocated_bytes=torch.cuda.max_memory_allocated(self.device), | |
| cuda_peak_reserved_bytes=torch.cuda.max_memory_reserved(self.device), | |
| ) | |
| return result | |
| def synchronize(self): | |
| if self.device is not None and str(self.device).startswith("cuda"): | |
| import torch | |
| torch.cuda.synchronize(self.device) | |
| def phase(self, name, **details): | |
| self.synchronize() | |
| start = time.monotonic() | |
| print(f"[prism] {name}: start", flush=True) | |
| event = {"phase": name, **details} | |
| try: | |
| yield | |
| self.synchronize() | |
| event["status"] = "ok" | |
| except BaseException as error: | |
| event.update(status="failed", error=f"{type(error).__name__}: {error}") | |
| raise | |
| finally: | |
| event.update(seconds=time.monotonic() - start, **self.memory()) | |
| self.events.append(event) | |
| print(f"[prism] {name}: {event['status']} ({event['seconds']:.2f}s)", flush=True) | |
| def wrap(self, owner, name, phase, validate=None, capture=None): | |
| original = getattr(owner, name) | |
| def measured(*args, **kwargs): | |
| with self.phase(phase): | |
| if capture is not None: | |
| capture(*args, **kwargs) | |
| result = original(*args, **kwargs) | |
| if validate is not None: | |
| validate(result) | |
| return result | |
| setattr(owner, name, measured) | |
| class DecodeLatentCapture: | |
| """Opt-in independent CPU copies of the exact two final VAE inputs.""" | |
| FORMAT = "prism-final-decoder-inputs-v1" | |
| NAMES = ("video_vae_input", "audio_vae_input") | |
| def __init__(self, metadata=None): | |
| self.metadata = dict(metadata or {}) | |
| self.tensors = {} | |
| def capture(self, name, tensor): | |
| import torch | |
| if name not in self.NAMES: | |
| raise ValueError(f"Unknown final decoder input: {name}") | |
| if name in self.tensors: | |
| raise RuntimeError(f"Final decoder input was captured more than once: {name}") | |
| if not torch.is_tensor(tensor) or tensor.is_meta: | |
| raise ValueError(f"Expected a materialized decoder input tensor: {name}") | |
| # copy=True matters when the input is already CPU: the decoder may later | |
| # mutate its own input. No dtype conversion, random calls or GPU cache. | |
| self.tensors[name] = tensor.detach().to(device="cpu", copy=True).contiguous() | |
| def contract(self): | |
| return {**self.metadata, "format": self.FORMAT, | |
| "capture_point": "Immediately before the public VAE decode call, before decoder-internal transforms", | |
| "video_vae_input_normalization": "Already denormalized using video_vae latents_mean/latents_std; do not denormalize again", | |
| "audio_vae_input_normalization": "Native final DAC diffusion latents; no extra normalization applied before decode", | |
| "video_layout": "B,C,T_latent,H_latent,W_latent", "audio_layout": "B,C,T_latent", | |
| "tensors": {name: {"shape": [int(n) for n in value.shape], | |
| "dtype": str(value.dtype).removeprefix("torch."), "device": "cpu"} | |
| for name, value in self.tensors.items()}} | |
| def save(self, path): | |
| from safetensors.torch import save_file | |
| path = Path(path) | |
| if path.suffix != ".safetensors": | |
| raise ValueError("Decoder inputs must be saved as Safetensors") | |
| if path.exists(): | |
| raise FileExistsError(f"Refusing to overwrite decoder inputs: {path}") | |
| if set(self.tensors) != set(self.NAMES): | |
| raise RuntimeError("Both final video and audio decoder inputs are required") | |
| contract = self.contract() | |
| metadata = {"format": self.FORMAT, "decoder_contract_json": json.dumps(contract, sort_keys=True, allow_nan=False)} | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| fd, temporary = tempfile.mkstemp(prefix=path.name + ".", suffix=".tmp", dir=path.parent) | |
| os.close(fd) | |
| temporary = Path(temporary) | |
| try: | |
| save_file(self.tensors, str(temporary), metadata=metadata) | |
| digest = hashlib.sha256() | |
| with temporary.open("rb") as stream: | |
| while chunk := stream.read(8 * 1024 * 1024): | |
| digest.update(chunk) | |
| os.fsync(stream.fileno()) | |
| size = temporary.stat().st_size | |
| # Same-directory hard-link publication is atomic and fails if a | |
| # concurrent writer already created the destination. Never replace. | |
| os.link(temporary, path) | |
| directory = os.open(path.parent, os.O_RDONLY) | |
| try: | |
| os.fsync(directory) | |
| finally: | |
| os.close(directory) | |
| finally: | |
| temporary.unlink(missing_ok=True) | |
| return {"path": str(path.resolve()), "sha256": digest.hexdigest(), "bytes": size, | |
| "format": self.FORMAT, "decoder_contract": contract} | |
| def _move_tensor(tensor, device, parameter): | |
| import torch | |
| moved = tensor.detach().to(device=device, non_blocking=False) | |
| return torch.nn.Parameter(moved, requires_grad=False) if parameter else moved | |
| class TensorStager: | |
| """Keep immutable CPU masters, temporarily bind device copies to a module. | |
| Restoring CPU references avoids copying unchanged weights back over PCIe. | |
| No dense dequantization cache is retained. Units must not overlap while active. | |
| """ | |
| def __init__(self, device, move_tensor=None): | |
| self.device = device | |
| self.move_tensor = move_tensor or _move_tensor | |
| self.active = {} | |
| self.transfers = 0 | |
| self.bytes_to_device = 0 | |
| def stage(self, module): | |
| identity = id(module) | |
| if identity in self.active: | |
| raise RuntimeError("A staging unit was entered twice without release") | |
| originals = [] | |
| copies = {} | |
| self.active[identity] = originals | |
| try: | |
| for owner in module.modules(): | |
| for parameter, registry in ((True, owner._parameters), (False, owner._buffers)): | |
| for name, tensor in list(registry.items()): | |
| if tensor is None: | |
| continue | |
| if getattr(tensor, "is_meta", False): | |
| raise RuntimeError(f"Cannot stage unmaterialized tensor {name}") | |
| if str(tensor.device) != "cpu": | |
| raise RuntimeError(f"Expected CPU master for {name}, found {tensor.device}") | |
| originals.append((registry, name, tensor)) | |
| key = (id(tensor), parameter) | |
| if key not in copies: | |
| copies[key] = self.move_tensor(tensor, self.device, parameter) | |
| self.bytes_to_device += tensor.numel() * tensor.element_size() | |
| registry[name] = copies[key] | |
| self.transfers += 1 | |
| except BaseException: | |
| self.release(module) | |
| raise | |
| def release(self, module): | |
| originals = self.active.pop(id(module), []) | |
| for registry, name, original in reversed(originals): | |
| registry[name] = original | |
| def release_all(self): | |
| for originals in list(self.active.values()): | |
| for registry, name, original in reversed(originals): | |
| registry[name] = original | |
| self.active.clear() | |
| def _stems(dit): | |
| if any(value is not None for value in dit._parameters.values()): | |
| raise RuntimeError("Unsupported DiT root parameters: update explicit stem staging") | |
| return [module for name, module in dit.named_children() if name != "blocks"] | |
| class OffloadedTransformer: | |
| """Pipeline-compatible adapter; .to(cuda) never moves the whole transformer.""" | |
| def __init__(self, model, device, telemetry=None, stager=None, validate_outputs=False): | |
| self.model = model | |
| self.device = device | |
| self.telemetry = telemetry | |
| self.stager = stager or TensorStager(device) | |
| self.handles = [] | |
| self.patches = [] | |
| self.validate_outputs = validate_outputs | |
| self.active_expert = None | |
| self.shared_active = False | |
| self.ready = False | |
| self.forward_calls = {"high": 0, "low": 0} | |
| self.block_calls = {"high_video": 0, "low_video": 0, "audio": 0, "a2v": 0, "v2a": 0} | |
| self.expert_switches = [] | |
| self.sparse_forward_calls = 0 | |
| self.sparse_kernel_calls = {"high": 0, "low": 0} | |
| self.static_sparse_kernel_calls = {"high": 0, "low": 0} | |
| self.dynamic_sparse_kernel_calls = {"high": 0, "low": 0} | |
| self.denoise_started = None | |
| self.shared_stems = _stems(model.audio_dit) + [model.dual_tower_bridge] | |
| self.expert_stems = {"high": _stems(model.video_dit), "low": _stems(model.video_dit_2)} | |
| for fused in model.fusion_blocks: | |
| for name, category in (("video_block", "high_video"), ("audio_block", "audio"), | |
| ("a2v_conditioner", "a2v"), ("v2a_conditioner", "v2a")): | |
| module = getattr(fused, name, None) | |
| if module is not None: | |
| self._hook(module, category) | |
| for module in model.remaining_video_blocks: | |
| self._hook(module, "high_video") | |
| for module in model.video_dit_2.blocks: | |
| self._hook(module, "low_video") | |
| def _hook(self, module, category): | |
| def before(owner, _args): | |
| self.stager.stage(owner) | |
| self.block_calls[category] += 1 | |
| if category.endswith("video") and getattr(owner.self_attn, "enable_bsa", False): | |
| self.sparse_forward_calls += 1 | |
| def after(owner, _args, _output): | |
| self.stager.release(owner) | |
| self.handles.append(module.register_forward_pre_hook(before)) | |
| self.handles.append(module.register_forward_hook(after, always_call=True)) | |
| def to(self, device): | |
| if str(device) == "cpu": | |
| if self.denoise_started is not None and self.telemetry is not None: | |
| self.telemetry.synchronize() | |
| self.telemetry.events.append({"phase": "denoising", "status": "ok", | |
| "seconds": time.monotonic() - self.denoise_started, | |
| **self.telemetry.memory()}) | |
| self.release() | |
| elif str(device) == str(self.device): | |
| self.ready = True | |
| if self.denoise_started is None: | |
| self.denoise_started = time.monotonic() | |
| else: | |
| raise ValueError(f"Adapter is configured for {self.device}, not {device}") | |
| return self | |
| def _activate(self, expert): | |
| if not self.shared_active: | |
| for module in self.shared_stems: | |
| self.stager.stage(module) | |
| self.shared_active = True | |
| if self.active_expert == expert: | |
| return | |
| if self.active_expert is not None: | |
| for module in self.expert_stems[self.active_expert]: | |
| self.stager.release(module) | |
| for module in self.expert_stems[expert]: | |
| self.stager.stage(module) | |
| self.active_expert = expert | |
| self.expert_switches.append(expert) | |
| print(f"[prism] active video expert: {expert}; transformer blocks stream from CPU", flush=True) | |
| def __call__(self, *args, **kwargs): | |
| if not self.ready: | |
| raise RuntimeError("Pipeline must enter transformer phase before forward") | |
| expert = "low" if kwargs.get("use_video_dit_2", False) else "high" | |
| try: | |
| self._activate(expert) | |
| self.forward_calls[expert] += 1 | |
| output = self.model(*args, **kwargs) | |
| if self.validate_outputs: | |
| if not isinstance(output, tuple) or len(output) != 2: | |
| raise ValueError("Official transformer must return both video and audio predictions") | |
| _assert_finite(output[0], "video noise prediction") | |
| _assert_finite(output[1], "audio noise prediction") | |
| return output | |
| except BaseException: | |
| self.release() | |
| raise | |
| def release(self): | |
| self.stager.release_all() | |
| self.active_expert = None | |
| self.shared_active = False | |
| self.ready = False | |
| self.denoise_started = None | |
| def instrument_sparse_attention(self, attention_module, dynamic_module=None): | |
| def instrument(owner, name, kind_counts): | |
| original = getattr(owner, name) | |
| def counted(*args, **kwargs): | |
| expert = self.active_expert | |
| if expert not in self.sparse_kernel_calls: | |
| raise RuntimeError("Sparse attention executed outside an active video expert") | |
| output = original(*args, **kwargs) | |
| kind_counts[expert] += 1 | |
| self.sparse_kernel_calls[expert] += 1 | |
| return output | |
| setattr(owner, name, counted) | |
| self.patches.append((owner, name, original, counted)) | |
| instrument(attention_module, "flash_attn_bsa_3d", self.static_sparse_kernel_calls) | |
| if dynamic_module is not None: | |
| instrument(dynamic_module, "flash_attn_bsa_3d_dynamic", self.dynamic_sparse_kernel_calls) | |
| def close(self): | |
| self.release() | |
| for handle in self.handles: | |
| handle.remove() | |
| self.handles.clear() | |
| for owner, name, original, replacement in self.patches: | |
| if getattr(owner, name) is replacement: | |
| setattr(owner, name, original) | |
| self.patches.clear() | |
| def receipt(self): | |
| return {"offload": "block", "forward_calls": dict(self.forward_calls), | |
| "block_calls": dict(self.block_calls), "expert_activations": list(self.expert_switches), | |
| "sparse_enabled_video_block_calls": self.sparse_forward_calls, | |
| "completed_sparse_kernel_calls": dict(self.sparse_kernel_calls), | |
| "completed_static_sparse_kernel_calls": dict(self.static_sparse_kernel_calls), | |
| "completed_dynamic_sparse_kernel_calls": dict(self.dynamic_sparse_kernel_calls), | |
| "staged_units": self.stager.transfers, "bytes_staged_to_device": self.stager.bytes_to_device} | |
| def _restore_runtime_tensors(model): | |
| import torch | |
| from hymm.models.modules.wan_video_dit import precompute_freqs_cis_3d | |
| from hymm.models.modules.wan_audio_dit import precompute_freqs_cis_1d, legacy_precompute_freqs_cis_1d | |
| with torch.device("cpu"): | |
| for dit in (model.video_dit, model.video_dit_2): | |
| head_dim = dit.dim // int(dit.config.num_heads) | |
| dit.freqs = precompute_freqs_cis_3d(head_dim) | |
| audio = model.audio_dit | |
| head_dim = audio.dim // int(audio.config.num_heads) | |
| if audio.vae_type == "dac": | |
| audio.freqs = precompute_freqs_cis_1d(head_dim) | |
| elif audio.vae_type == "oobleck": | |
| audio.freqs = legacy_precompute_freqs_cis_1d(head_dim) | |
| else: | |
| raise ValueError(f"Unsupported audio VAE: {audio.vae_type}") | |
| rotary = model.dual_tower_bridge.rotary | |
| rotary.inv_freq = 1.0 / (rotary.base ** (torch.arange(0, rotary.dim, 2).float() / rotary.dim)) | |
| rotary.original_inv_freq = rotary.inv_freq | |
| for name, tensor in list(model.named_parameters()) + list(model.named_buffers()): | |
| if tensor.is_meta: | |
| raise RuntimeError(f"Unmaterialized model tensor: {name}") | |
| for dit in (model.video_dit, model.video_dit_2, model.audio_dit): | |
| if any(tensor.is_meta for tensor in dit.freqs): | |
| raise RuntimeError("Unmaterialized rotary frequency tuple") | |
| def create_meta_transformer(base_dir, boundary_ratio=0.9): | |
| """Small allocation-only integration probe; prepared_source must run first.""" | |
| import torch | |
| from hymm.models.modules.mova import MOVABridge | |
| from hymm.models.modules.wan_video_dit import WanModel | |
| from hymm.models.modules.wan_audio_dit import WanAudioModel | |
| from hymm.models.modules.interactionv2 import DualTowerConditionalBridge | |
| base_dir = Path(base_dir) | |
| with torch.device("meta"): | |
| video = WanModel.from_config(read_json(base_dir / "video_dit/config.json")) | |
| video_2 = WanModel.from_config(read_json(base_dir / "video_dit_2/config.json")) | |
| audio = WanAudioModel.from_config(read_json(base_dir / "audio_dit/config.json")) | |
| bridge = DualTowerConditionalBridge.from_config(read_json(base_dir / "dual_tower_bridge/config.json")) | |
| model = MOVABridge(video, video_2, audio, bridge, boundary_ratio=boundary_ratio) | |
| # Nonpersistent rotary buffers are absent from the checkpoint. Restore before | |
| # the strict loader's no-registered-meta check, then check all tensors again. | |
| rotary = model.dual_tower_bridge.rotary | |
| rotary.inv_freq = 1.0 / (rotary.base ** (torch.arange(0, rotary.dim, 2, device="cpu").float() / rotary.dim)) | |
| rotary.original_inv_freq = rotary.inv_freq | |
| return model | |
| def validate_precision(precision, source_checkpoint): | |
| if precision not in ("q6", "original-bf16"): | |
| raise ValueError("Precision must be q6 or original-bf16") | |
| if precision == "original-bf16" and source_checkpoint is None: | |
| raise ValueError("original-bf16 requires an explicit --source-checkpoint") | |
| if precision == "q6" and source_checkpoint is not None: | |
| raise ValueError("--source-checkpoint is only used with --precision original-bf16") | |
| def validate_attention_options(mode, ivpq_dynamic_block=False): | |
| if mode not in ("sparse", "dense", "high-only"): | |
| raise ValueError("Video attention must be sparse, dense, or high-only") | |
| if ivpq_dynamic_block and mode == "dense": | |
| raise ValueError("--ivpq-dynamic-block requires sparse or high-only video attention") | |
| def configure_video_attention(model, mode, ivpq_dynamic_block=False): | |
| """Select official BSA scope and optional IVPQ; never replace attention math.""" | |
| validate_attention_options(mode, ivpq_dynamic_block) | |
| # This iterator deliberately ignores the previous expert scope. Clear stale | |
| # dynamic flags on BOTH experts, including a low expert about to be skipped. | |
| for attention in model._video_self_attns(): | |
| attention.enable_ivpq_dynamic_block = False | |
| attention.enable_penalty_dynamic_block = False | |
| attention.enable_layer_adaptive_dynamic_block = False | |
| params = {"sparsity": 0.75, "cdf_threshold": 0.20, | |
| "chunk_3d_shape_q": [4, 4, 4], "chunk_3d_shape_k": [4, 4, 4]} | |
| if ivpq_dynamic_block: | |
| params.update(dynamic_block_lambda_a=0.5, dynamic_block_tau_128=0.25, | |
| dynamic_block_lambda_128=1.0) | |
| if mode == "high-only": | |
| # The upstream scope setter only changes a flag. Clear both experts | |
| # first so a previously enabled low expert cannot retain sparse BSA. | |
| model.set_sparse_high_noise_only(False) | |
| model.configure_bsa(False, bsa_params=None, enable_bsa_v2a=False, bsa_params_v2a=None) | |
| model.set_sparse_high_noise_only(mode == "high-only") | |
| model.configure_bsa(enable_bsa=mode != "dense", bsa_params=params if mode != "dense" else None, | |
| enable_bsa_v2a=False, bsa_params_v2a=None) | |
| if ivpq_dynamic_block: | |
| model.configure_ivpq_dynamic_block(True) | |
| sparse_kind = "ivpq-bsa" if ivpq_dynamic_block else "static-bsa" | |
| return {"policy": mode, "high_expert": "dense" if mode == "dense" else sparse_kind, | |
| "low_expert": sparse_kind if mode == "sparse" else "dense", | |
| "video_both_experts": mode == "sparse", "sparsity": 0.75, "cdf_threshold": 0.20, | |
| "block_shape": None if ivpq_dynamic_block else [4, 4, 4], | |
| "dynamic_blocks": ivpq_dynamic_block, "bridge_v2a": False, | |
| "ivpq": {"enabled": ivpq_dynamic_block, "lambda_a": 0.5 if ivpq_dynamic_block else None, | |
| "tau_128": 0.25 if ivpq_dynamic_block else None, | |
| "macro_zone_size": 8 if ivpq_dynamic_block else None, | |
| "penalty_enabled": False, "penalty_lambda_128": 1.0 if ivpq_dynamic_block else None, | |
| "layer_adaptive_enabled": False}} | |
| def validate_execution(transformer, video_attention, ivpq_dynamic_block=False): | |
| """Require both experts and only the sparse kernels selected for this run.""" | |
| validate_attention_options(video_attention, ivpq_dynamic_block) | |
| if not all(transformer.forward_calls[expert] > 0 for expert in ("high", "low")): | |
| raise RuntimeError("Generation did not exercise both video experts") | |
| expected_sparse = {"high": video_attention != "dense", "low": video_attention == "sparse"} | |
| for expert, enabled in expected_sparse.items(): | |
| calls = transformer.sparse_kernel_calls[expert] | |
| static = transformer.static_sparse_kernel_calls[expert] | |
| dynamic = transformer.dynamic_sparse_kernel_calls[expert] | |
| if calls != static + dynamic: | |
| raise RuntimeError(f"Official sparse attention kernel counters are inconsistent for the {expert} expert") | |
| if enabled and calls <= 0: | |
| raise RuntimeError(f"Official sparse attention kernel was not exercised by the {expert} expert") | |
| if not enabled and calls != 0: | |
| raise RuntimeError(f"Unexpected sparse attention kernel execution by the dense {expert} expert") | |
| if enabled and ((ivpq_dynamic_block and static) or (not ivpq_dynamic_block and dynamic)): | |
| raise RuntimeError(f"Unexpected sparse attention kernel mode for the {expert} expert") | |
| if video_attention != "dense" and not transformer.sparse_forward_calls: | |
| raise RuntimeError("No video block with official sparse attention was executed") | |
| if video_attention == "dense" and transformer.sparse_forward_calls: | |
| raise RuntimeError("Unexpected sparse-enabled video block in dense attention mode") | |
| if sum(transformer.sparse_kernel_calls.values()) != transformer.sparse_forward_calls: | |
| raise RuntimeError("Official sparse attention kernel count differs from sparse-enabled video block calls") | |
| class VAETilingPolicy: | |
| """Apply phase-specific tiling, including when a pipeline is called again.""" | |
| def __init__(self, vae, mode="both", tile_size=256, tile_stride=192, latent_capture=None): | |
| if mode not in ("both", "decode", "off"): | |
| raise ValueError("VAE tiling must be both, decode, or off") | |
| if tile_size <= 0 or tile_stride <= 0 or tile_stride >= tile_size or tile_size % 8 or tile_stride % 8: | |
| raise ValueError("VAE tile size/stride must be multiples of 8 with 0 < stride < size") | |
| self.vae, self.mode = vae, mode | |
| self.tile_size, self.tile_stride = tile_size, tile_stride | |
| self.encode_enabled, self.decode_enabled = mode == "both", mode != "off" | |
| self.latent_capture = latent_capture | |
| def _apply(self, enabled): | |
| if enabled: | |
| self.vae.enable_tiling(tile_sample_min_height=self.tile_size, tile_sample_min_width=self.tile_size, | |
| tile_sample_stride_height=self.tile_stride, tile_sample_stride_width=self.tile_stride) | |
| else: | |
| self.vae.disable_tiling() | |
| def before_encode(self, *args, **kwargs): | |
| self._apply(self.encode_enabled) | |
| def before_decode(self, z, *args, **kwargs): | |
| self._apply(self.decode_enabled) | |
| if self.latent_capture is not None: | |
| self.latent_capture.metadata.update( | |
| vae_tiling={"enabled": bool(self.vae.use_tiling), "tile_size": self.tile_size, | |
| "stride": self.tile_stride}, | |
| conditioning_vae_tiling=self.encode_enabled, vae_tiling_mode=self.mode) | |
| self.latent_capture.capture("video_vae_input", z) | |
| def receipt(self): | |
| return {"mode": self.mode, "size": self.tile_size, "stride": self.tile_stride, | |
| "encode": self.encode_enabled, "decode": self.decode_enabled} | |
| def _build_transformer(base_dir, manifest, boundary_ratio, precision="q6", source_checkpoint=None, | |
| video_attention="sparse", ivpq_dynamic_block=False): | |
| validate_precision(precision, source_checkpoint) | |
| validate_attention_options(video_attention, ivpq_dynamic_block) | |
| model = create_meta_transformer(base_dir, boundary_ratio) | |
| if precision == "original-bf16": | |
| # Select this loader before any Q6 replacement/allocation. This diagnostic | |
| # compares the original weights using otherwise identical inference code. | |
| from .baseline import load_original_bf16 | |
| model, receipt = load_original_bf16(model, source_checkpoint, manifest, strict=True) | |
| if receipt.get("runtime_quantized_modules") != 0 or not receipt.get("diagnostic_baseline"): | |
| raise RuntimeError("Original-BF16 loader did not certify zero quantized runtime modules") | |
| else: | |
| from .quant_loader import load_quantized_model | |
| model, receipt = load_quantized_model(model, manifest, device="cpu", strict=True) | |
| receipt.update(effective_runtime_profile=receipt["profile"], | |
| runtime_quantized_modules=receipt["quantized_modules"], diagnostic_baseline=False) | |
| receipt["runtime_precision"] = precision | |
| _restore_runtime_tensors(model) | |
| for dit in (model.video_dit, model.video_dit_2, model.audio_dit): | |
| dit.time_embedding.float() | |
| dit.time_projection.float() | |
| model.eval().requires_grad_(False) | |
| receipt["time_modules_dtype"] = "float32" | |
| receipt["bsa"] = configure_video_attention(model, video_attention, ivpq_dynamic_block) | |
| receipt["video_attention"] = video_attention | |
| return model, receipt | |
| def _assert_finite(tensor, label): | |
| import torch | |
| if not torch.is_tensor(tensor) or not bool(torch.isfinite(tensor).all().item()): | |
| raise ValueError(f"Non-finite or invalid {label}") | |
| class _CPUTextEncoder: | |
| """Keep the frozen encoder on CPU despite the upstream phase-level .to().""" | |
| def __init__(self, encoder): | |
| self.encoder = encoder | |
| def dtype(self): | |
| return self.encoder.dtype | |
| def to(self, _device): | |
| return self | |
| def __call__(self, *args, **kwargs): | |
| return self.encoder(*args, **kwargs) | |
| def build_pipeline(base_dir, manifest, device, telemetry, text_device="cpu", vendor_dir=None, | |
| dense_attention="sdpa", tile_size=256, tile_stride=192, | |
| precision="q6", source_checkpoint=None, latent_capture=None, | |
| vae_tiling="both", video_attention="sparse", ivpq_dynamic_block=False): | |
| """Load Q6 or diagnostic original preview weights; frozen assets stay shared.""" | |
| validate_precision(precision, source_checkpoint) | |
| validate_attention_options(video_attention, ivpq_dynamic_block) | |
| import torch | |
| vendor = prepared_source(vendor_dir or PROJECT_ROOT / "vendor/prism") | |
| guidance_patch = prepared_guidance_source(vendor) | |
| from diffusers.models.autoencoders import AutoencoderKLWan | |
| from transformers import T5TokenizerFast, UMT5EncoderModel | |
| from hymm.models.modules.dac_vae import DAC | |
| from hymm.models.modules import wan_video_dit | |
| from hymm.models.modules.block_sparse_attention import dynamic_block_attention | |
| from hymm.diffusion.schedulers.flow_match_pair import FlowMatchPairScheduler | |
| from hymm.diffusion.pipelines.mova_pipeline import MOVAPipeline | |
| if dense_attention == "sdpa": | |
| wan_video_dit.FLASH_ATTN_3_AVAILABLE = False | |
| wan_video_dit.FLASH_ATTN_2_AVAILABLE = False | |
| wan_video_dit.SAGE_ATTN_AVAILABLE = False | |
| elif dense_attention != "local": | |
| raise ValueError("Dense attention must be sdpa or local") | |
| base_dir = Path(base_dir).resolve() | |
| index = read_json(base_dir / "model_index.json") | |
| if index.get("audio_vae_type", "dac") != "dac": | |
| raise ValueError("This runner currently supports the official DAC asset set") | |
| boundary = float(index.get("boundary_ratio", 0.9)) | |
| if not 0 < boundary < 1: | |
| raise ValueError("Invalid expert boundary ratio") | |
| load_phase = "load_original_bf16_transformer" if precision == "original-bf16" else "load_quantized_transformer" | |
| with telemetry.phase(load_phase): | |
| model, load_receipt = _build_transformer(base_dir, manifest, boundary, precision, source_checkpoint, | |
| video_attention=video_attention, ivpq_dynamic_block=ivpq_dynamic_block) | |
| with telemetry.phase("load_frozen_components"): | |
| local = {"local_files_only": True, "use_safetensors": True, "low_cpu_mem_usage": True} | |
| video_vae = AutoencoderKLWan.from_pretrained(base_dir / "video_vae", torch_dtype=torch.bfloat16, **local) | |
| # DAC's weight-norm construction is not assumed compatible with the | |
| # generic meta loader. This small component loads normally on CPU. | |
| audio_vae = DAC.from_pretrained(base_dir / "audio_vae", torch_dtype=torch.float32, | |
| local_files_only=True, use_safetensors=True, low_cpu_mem_usage=False) | |
| encoder = UMT5EncoderModel.from_pretrained(base_dir / "text_encoder", torch_dtype=torch.bfloat16, **local) | |
| tokenizer = T5TokenizerFast.from_pretrained(base_dir / "tokenizer", local_files_only=True) | |
| scheduler = FlowMatchPairScheduler.from_pretrained(base_dir / "scheduler", local_files_only=True) | |
| for component in (video_vae, audio_vae, encoder): | |
| component.eval().requires_grad_(False) | |
| tiling_policy = VAETilingPolicy(video_vae, vae_tiling, tile_size, tile_stride, latent_capture) | |
| tiling_policy.before_encode() | |
| transformer = OffloadedTransformer(model, device, telemetry, validate_outputs=True) | |
| transformer.instrument_sparse_attention(wan_video_dit, dynamic_block_attention) | |
| pipeline = MOVAPipeline(transformer, video_vae, audio_vae, | |
| _CPUTextEncoder(encoder) if text_device == "cpu" else encoder, | |
| tokenizer, scheduler, boundary_ratio=boundary, device=device) | |
| pipeline.enable_cpu_offload(gpu_device=device) | |
| if text_device == "cpu": | |
| original = pipeline._get_t5_prompt_embeds | |
| def cpu_embeddings(*args, **kwargs): | |
| kwargs["device"] = torch.device("cpu") | |
| return original(*args, **kwargs).to(device=device) | |
| pipeline._get_t5_prompt_embeds = cpu_embeddings | |
| elif text_device != "cuda": | |
| raise ValueError("Text device must be cpu or cuda") | |
| telemetry.wrap(pipeline, "_get_t5_prompt_embeds", "text_encoding") | |
| telemetry.wrap(pipeline, "prepare_latents", "condition_encoding", capture=tiling_policy.before_encode) | |
| telemetry.wrap(video_vae, "decode", "video_decode", lambda output: _assert_finite(output.sample, "decoded video"), | |
| capture=tiling_policy.before_decode) | |
| telemetry.wrap(audio_vae, "decode", "audio_decode", lambda output: _assert_finite(output, "decoded audio"), | |
| capture=(lambda z, *args, **kwargs: latent_capture.capture("audio_vae_input", z)) if latent_capture else None) | |
| load_receipt.update(source_dir=str(vendor), text_device=text_device, dense_attention=dense_attention, | |
| audio_vae_dtype="float32", boundary_ratio=boundary, | |
| vae_tiling=tiling_policy.receipt(), guidance_source_patch=guidance_patch) | |
| return pipeline, transformer, load_receipt | |
| def cleanup_pipeline(pipeline, transformer): | |
| transformer.close() | |
| for name in ("text_encoder", "video_vae", "audio_vae"): | |
| getattr(pipeline, name).to("cpu") | |
| def _probe_media(path, ffprobe="ffprobe"): | |
| result = subprocess.run([ffprobe, "-v", "error", "-count_frames", "-show_streams", "-show_format", | |
| "-of", "json", str(path)], check=True, capture_output=True, text=True) | |
| return json.loads(result.stdout) | |
| def validate_media_probe(probe, width, height, frames, duration): | |
| streams = probe.get("streams", []) | |
| video = [stream for stream in streams if stream.get("codec_type") == "video"] | |
| audio = [stream for stream in streams if stream.get("codec_type") == "audio"] | |
| if len(video) != 1 or len(audio) != 1: | |
| raise ValueError("Output must contain exactly one video and one generated audio stream") | |
| video, audio = video[0], audio[0] | |
| if (video.get("width"), video.get("height")) != (width, height): | |
| raise ValueError("Encoded dimensions differ from requested dimensions") | |
| if int(video.get("nb_read_frames", video.get("nb_frames", -1))) != frames: | |
| raise ValueError("Encoded frame count differs from generated frame count") | |
| if video.get("codec_name") != "h264" or audio.get("codec_name") != "aac": | |
| raise ValueError("Expected H.264 video and AAC generated audio") | |
| actual_duration = float(probe.get("format", {}).get("duration", 0)) | |
| if not math.isfinite(actual_duration) or abs(actual_duration - duration) > 0.15: | |
| raise ValueError(f"Unexpected media duration: {actual_duration}, expected {duration}") | |
| return {"width": width, "height": height, "frames": frames, "duration_seconds": actual_duration, | |
| "video_codec": video["codec_name"], "audio_codec": audio["codec_name"], | |
| "audio_sample_rate": int(audio["sample_rate"]), "audio_channels": int(audio["channels"])} | |
| def save_media(frames, audio, output, fps, sample_rate, expected_frames, width, height, | |
| ffmpeg="ffmpeg", ffprobe="ffprobe"): | |
| """Encode, mux, decode-check, and atomically publish both MP4 and source WAV.""" | |
| import numpy as np | |
| output = Path(output).resolve() | |
| wav_output = output.with_suffix(".wav") | |
| if output.exists() or wav_output.exists(): | |
| raise FileExistsError(f"Refusing to overwrite output or generated audio: {output}") | |
| if len(frames) != expected_frames: | |
| raise ValueError(f"Generated {len(frames)} frames; expected {expected_frames}") | |
| if hasattr(audio, "detach"): | |
| audio = audio.detach().float().cpu().numpy() | |
| audio = np.asarray(audio, dtype=np.float32) | |
| if audio.ndim == 1: | |
| audio = audio[None, :] | |
| if audio.ndim != 2 or audio.shape[0] not in (1, 2): | |
| raise ValueError(f"Expected mono or stereo generated audio, got {audio.shape}") | |
| sample_count = int(sample_rate * expected_frames / fps) | |
| if audio.shape[1] < sample_count: | |
| raise ValueError("Generated audio is shorter than the video") | |
| audio = audio[:, :sample_count] | |
| if not np.isfinite(audio).all(): | |
| raise ValueError("Generated audio contains NaN or infinity") | |
| rms = float(np.sqrt(np.mean(np.square(audio.astype(np.float64))))) | |
| if rms <= 1e-8: | |
| raise ValueError("Generated audio is silent") | |
| metrics = {"source_audio_samples": sample_count, "source_audio_rms": rms, | |
| "source_audio_peak": float(np.max(np.abs(audio))), | |
| "source_audio_clipped_fraction": float(np.mean(np.abs(audio) > 1))} | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| with tempfile.TemporaryDirectory(prefix="prism-output-", dir=output.parent) as temporary: | |
| tmp = Path(temporary) | |
| video_path, audio_path, mux_path = tmp / "video.mp4", tmp / "audio.wav", tmp / "output.mp4" | |
| with wave.open(str(audio_path), "wb") as stream: | |
| stream.setnchannels(audio.shape[0]) | |
| stream.setsampwidth(2) | |
| stream.setframerate(int(sample_rate)) | |
| stream.writeframes((np.clip(audio, -1, 1).T * 32767).astype("<i2").tobytes()) | |
| command = [ffmpeg, "-hide_banner", "-loglevel", "error", "-nostdin", "-y", "-f", "rawvideo", | |
| "-pixel_format", "rgb24", "-video_size", f"{width}x{height}", "-framerate", str(fps), | |
| "-i", "pipe:0", "-an", "-c:v", "libx264", "-crf", "18", "-pix_fmt", "yuv420p", | |
| "-movflags", "+faststart", str(video_path)] | |
| differences = [] | |
| previous = None | |
| with (tmp / "encode.log").open("wb") as stderr: | |
| process = subprocess.Popen(command, stdin=subprocess.PIPE, stdout=subprocess.DEVNULL, stderr=stderr) | |
| try: | |
| for frame in frames: | |
| array = np.asarray(frame.convert("RGB") if hasattr(frame, "convert") else frame) | |
| if array.shape != (height, width, 3) or array.dtype != np.uint8: | |
| raise ValueError(f"Invalid generated video frame: {array.shape}, {array.dtype}") | |
| small = array[::8, ::8].astype(np.float32) | |
| if previous is not None: | |
| differences.append(float(np.mean(np.abs(small - previous)))) | |
| previous = small | |
| process.stdin.write(array.tobytes()) | |
| process.stdin.close() | |
| return_code = process.wait() | |
| if return_code: | |
| raise RuntimeError((tmp / "encode.log").read_text(errors="replace")[-2000:]) | |
| except BaseException: | |
| if process.stdin and not process.stdin.closed: | |
| with contextlib.suppress(BrokenPipeError): | |
| process.stdin.close() | |
| if process.poll() is None: | |
| process.terminate() | |
| process.wait() | |
| raise | |
| metrics["mean_sampled_frame_difference"] = float(np.mean(differences)) | |
| metrics["identical_sampled_frame_pairs"] = sum(value == 0 for value in differences) | |
| if metrics["mean_sampled_frame_difference"] <= 0: | |
| raise ValueError("Generated frames are identical at all sampled positions") | |
| result = subprocess.run([ffmpeg, "-hide_banner", "-loglevel", "error", "-nostdin", "-y", | |
| "-i", str(video_path), "-i", str(audio_path), "-map", "0:v:0", "-map", "1:a:0", | |
| "-c:v", "copy", "-c:a", "aac", "-b:a", "192k", | |
| "-movflags", "+faststart", str(mux_path)], capture_output=True, text=True) | |
| if result.returncode: | |
| raise RuntimeError(f"Audio/video mux failed: {result.stderr[-2000:]}") | |
| probe = _probe_media(mux_path, ffprobe) | |
| metrics.update(validate_media_probe(probe, width, height, expected_frames, expected_frames / fps)) | |
| decoded = subprocess.run([ffmpeg, "-v", "error", "-i", str(mux_path), "-map", "0:a:0", | |
| "-f", "f32le", "-acodec", "pcm_f32le", "pipe:1"], check=True, | |
| capture_output=True) | |
| decoded_audio = np.frombuffer(decoded.stdout, dtype="<f4") | |
| if decoded_audio.size == 0 or not np.isfinite(decoded_audio).all(): | |
| raise ValueError("Muxed audio failed decode/finite check") | |
| decoded_rms = float(np.sqrt(np.mean(decoded_audio.astype(np.float64) ** 2))) | |
| if decoded_rms <= 1e-8: | |
| raise ValueError("Muxed audio is silent") | |
| metrics["decoded_audio_rms"] = decoded_rms | |
| audio_path.replace(wav_output) | |
| mux_path.replace(output) | |
| metrics.update(output=str(output), generated_audio=str(wav_output)) | |
| return metrics | |