Prism-Q6 / loader /prism_quant /runtime.py
atomtanstudio's picture
Add model card, Q6 manifest, loader package and conversion receipts
18e823d verified
Raw History Blame Contribute Delete
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)
@contextlib.contextmanager
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)
@functools.wraps(original)
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)
@functools.wraps(original)
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
@property
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