Instructions to use ApacheOne/Wan2.2-Animate-2-14B-OrbitQuant-W4A4 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use ApacheOne/Wan2.2-Animate-2-14B-OrbitQuant-W4A4 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline from diffusers.utils import load_image, export_to_video # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("ApacheOne/Wan2.2-Animate-2-14B-OrbitQuant-W4A4", dtype=torch.bfloat16, device_map="cuda") pipe.to("cuda") prompt = "A man with short gray hair plays a red electric guitar." image = load_image( "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/guitar-man.png" ) output = pipe(image=image, prompt=prompt).frames[0] export_to_video(output, "output.mp4") - Notebooks
- Google Colab
- Kaggle
| from __future__ import annotations | |
| import gc | |
| import inspect | |
| import json | |
| import sys | |
| from pathlib import Path | |
| from typing import Any | |
| import torch | |
| import yaml | |
| from safetensors import safe_open | |
| from .packed_linear import OrbitQuantPackedLinear, OrbitQuantW4A4Engine | |
| from .rotation_bank import RotationBank | |
| def _weight_map(root: Path) -> dict[str, str]: | |
| indexes = sorted(root.glob('*.safetensors.index.json')) | |
| if indexes: | |
| return dict(json.loads(indexes[0].read_text())['weight_map']) | |
| out = {} | |
| for p in sorted(root.glob('*.safetensors')): | |
| if p.name == 'orbitquant_rotations.safetensors': | |
| continue | |
| with safe_open(str(p), framework='pt', device='cpu') as f: | |
| for k in f.keys(): | |
| out[k] = p.name | |
| return out | |
| def _resolve_parent(module: torch.nn.Module, path: str): | |
| parts = path.split('.') | |
| obj: Any = module | |
| for part in parts[:-1]: | |
| if part.isdigit(): | |
| obj = obj[int(part)] | |
| else: | |
| obj = getattr(obj, part) | |
| return obj, parts[-1] | |
| def _get_submodule(module: torch.nn.Module, path: str): | |
| obj: Any = module | |
| for part in path.split('.'): | |
| obj = obj[int(part)] if part.isdigit() else getattr(obj, part) | |
| return obj | |
| def _set_submodule(module: torch.nn.Module, path: str, value: torch.nn.Module): | |
| parent, leaf = _resolve_parent(module, path) | |
| if leaf.isdigit(): | |
| parent[int(leaf)] = value | |
| else: | |
| setattr(parent, leaf, value) | |
| def _assign_tensor(module: torch.nn.Module, key: str, value: torch.Tensor): | |
| parent, leaf = _resolve_parent(module, key) | |
| if leaf in getattr(parent, '_parameters', {}): | |
| old = parent._parameters[leaf] | |
| req = bool(old.requires_grad) if old is not None else False | |
| if old is not None and tuple(old.shape) != tuple(value.shape): | |
| raise RuntimeError(f'shape mismatch for {key}: model={tuple(old.shape)} ckpt={tuple(value.shape)}') | |
| parent._parameters[leaf] = torch.nn.Parameter(value.contiguous(), requires_grad=req) | |
| return | |
| if leaf in getattr(parent, '_buffers', {}): | |
| old = parent._buffers[leaf] | |
| if old is not None and tuple(old.shape) != tuple(value.shape): | |
| raise RuntimeError(f'buffer shape mismatch for {key}') | |
| parent._buffers[leaf] = value.contiguous() | |
| return | |
| attr = getattr(parent, leaf, None) | |
| if torch.is_tensor(attr): | |
| setattr(parent, leaf, value.contiguous()) | |
| return | |
| raise KeyError(f'cannot assign checkpoint tensor to model path: {key}') | |
| def _official_ctor_kwargs(config: dict, cls) -> dict: | |
| aliases = { | |
| 'patch_size': config.get('patch_size', (1, 2, 2)), | |
| 'text_len': config.get('text_len', 512), | |
| 'in_dim': config.get('in_dim', config.get('in_channels', 36)), | |
| 'dim': config.get('dim', 5120), | |
| 'ffn_dim': config.get('ffn_dim', 13824), | |
| 'freq_dim': config.get('freq_dim', 256), | |
| 'text_dim': config.get('text_dim', 4096), | |
| 'out_dim': config.get('out_dim', config.get('out_channels', 16)), | |
| 'num_heads': config.get('num_heads', config.get('num_attention_heads', 40)), | |
| 'num_layers': config.get('num_layers', 40), | |
| 'window_size': tuple(config.get('window_size', (-1, -1))), | |
| 'qk_norm': config.get('qk_norm', True), | |
| 'cross_attn_norm': config.get('cross_attn_norm', True), | |
| 'eps': float(config.get('eps', 1e-6)), | |
| 'use_img_emb': config.get('use_img_emb', True), | |
| 'refer_offset_t': config.get('refer_offset_t', 1), | |
| 'refer_offset_h': config.get('refer_offset_h', 0), | |
| 'refer_offset_w': config.get('refer_offset_w', -1), | |
| 'refer_stride': config.get('refer_stride', 1), | |
| 'sparse_type': config.get('sparse_type', 0), | |
| 'use_context_parallel': False, | |
| 'log_scale': float(config.get('log_scale', 0.0)), | |
| } | |
| sig = inspect.signature(cls.__init__) | |
| return {k: v for k, v in aliases.items() if k in sig.parameters} | |
| def build_packed_official_transformer( | |
| official_repo_dir: str | Path, | |
| packed_dir: str | Path, | |
| *, | |
| config_path: str | Path | None = None, | |
| ) -> tuple[torch.nn.Module, dict]: | |
| official_repo_dir = Path(official_repo_dir).resolve() | |
| pdir = Path(packed_dir).resolve() | |
| manifest = json.loads((pdir / 'packed_manifest.json').read_text()) | |
| if manifest['target_count'] != 480: | |
| raise RuntimeError('packed artifact does not contain exactly 480 targets') | |
| if manifest.get('format') != 'OrbitQuant_WanAnimate2_direct_source_packed_nonuniform_W4A4_v3': | |
| raise RuntimeError(f"unsupported packed artifact format: {manifest.get('format')}") | |
| if manifest.get('weight_storage_layout') != 'K_half_by_N': | |
| raise RuntimeError('packed artifact is not in the GEMM-native [K/2,N] weight layout') | |
| if manifest.get('passthrough_count') != 823 or manifest.get('full_transformer_tensor_count') != 1303: | |
| raise RuntimeError( | |
| 'packed artifact is not the exact Wan-Animate-2 1303-tensor runtime inventory ' | |
| f"(passthrough={manifest.get('passthrough_count')}, total={manifest.get('full_transformer_tensor_count')})" | |
| ) | |
| config_file = Path(config_path) if config_path else (pdir / 'config.json') | |
| if not config_file.is_file(): | |
| raise FileNotFoundError(f'packed transformer config missing: {config_file}') | |
| config = json.loads(config_file.read_text()) | |
| # The official distilled runtime carries inference-critical transformer | |
| # fields (notably log_scale=-1.3) in its YAML. Diffusers' refactor config | |
| # did not historically preserve every such field, so the official source | |
| # config is authoritative for the official source model constructor. | |
| distill_yaml = official_repo_dir / 'infer' / 'wan_animate_2_distillation.yaml' | |
| if distill_yaml.is_file(): | |
| official_cfg = yaml.safe_load(distill_yaml.read_text()) | |
| for key, value in dict(official_cfg.get('model', {}).get('transformer', {})).items(): | |
| if key != 'type': | |
| config[key] = value | |
| if str(official_repo_dir) not in sys.path: | |
| sys.path.insert(0, str(official_repo_dir)) | |
| from wanxiang.wanxiang_animate_2_arch import WanxiangAnimate2Transformer | |
| kwargs = _official_ctor_kwargs(config, WanxiangAnimate2Transformer) | |
| with torch.device('meta'): | |
| model = WanxiangAnimate2Transformer(**kwargs) | |
| model.eval().requires_grad_(False) | |
| bank = RotationBank.load(pdir / manifest['rotation_bank']) | |
| engine = OrbitQuantW4A4Engine(bank) | |
| # Keep the activation pack cache scoped to one transformer block. Q/K/V | |
| # inside a block can share the same A4 pack, but retaining those uint4 | |
| # activations across blocks only wastes VRAM (the next block consumes a | |
| # different tensor). Counters remain cumulative when the entries clear. | |
| cache_clear_handles = [] | |
| for block in model.blocks: | |
| cache_clear_handles.append( | |
| block.register_forward_hook( | |
| lambda _module, _args, output, _engine=engine: (_engine.activation_cache.clear(), output)[1] | |
| ) | |
| ) | |
| # Replace the 480 dense Linear weights before materializing the checkpoint. | |
| target_stored = set(manifest['targets']) | |
| replaced = [] | |
| for stored_key, info in manifest['targets'].items(): | |
| official_weight = info['official_key'] | |
| module_path = official_weight[:-len('.weight')] | |
| old = _get_submodule(model, module_path) | |
| if not isinstance(old, torch.nn.Linear): | |
| raise RuntimeError(f'target is not nn.Linear in official model: {module_path} -> {type(old)}') | |
| new = OrbitQuantPackedLinear(old.in_features, old.out_features, old.bias is not None, engine) | |
| _set_submodule(model, module_path, new) | |
| replaced.append(module_path) | |
| if len(set(replaced)) != 480: | |
| raise RuntimeError(f'replaced {len(set(replaced))}, expected 480') | |
| # Load the packed uint4 weights/scales into those modules (CPU residency initially). | |
| by_shard = {} | |
| for stored_key, info in manifest['targets'].items(): | |
| by_shard.setdefault(info['shard'], []).append((stored_key, info)) | |
| for shard, items in sorted(by_shard.items()): | |
| with safe_open(str(pdir / shard), framework='pt', device='cpu') as sf: | |
| for _stored_key, info in items: | |
| module_path = info['official_key'][:-len('.weight')] | |
| mod = _get_submodule(model, module_path) | |
| mod.set_packed(sf.get_tensor(info['packed_tensor']), sf.get_tensor(info['scale_tensor'])) | |
| # Stream every non-target tensor from the self-contained packed runtime | |
| # artifact. The original BF16 transformer shards are no longer required after the | |
| # direct packing stage. | |
| passthrough = dict(manifest.get('passthrough_tensors', {})) | |
| if not passthrough: | |
| raise RuntimeError( | |
| 'packed artifact has no passthrough tensors; rebuild it with this runtime so the transformer is self-contained' | |
| ) | |
| grouped = {} | |
| for stored_key, info in passthrough.items(): | |
| grouped.setdefault(info['shard'], []).append((stored_key, info)) | |
| assigned = 0 | |
| for fname, items in sorted(grouped.items()): | |
| with safe_open(str(pdir / fname), framework='pt', device='cpu') as sf: | |
| for stored_key, info in sorted(items): | |
| _assign_tensor(model, info['official_key'], sf.get_tensor(info['tensor'])) | |
| assigned += 1 | |
| gc.collect() | |
| skipped_w4 = len(target_stored) | |
| meta_params = [n for n, p in model.named_parameters() if p.is_meta] | |
| meta_buffers = [n for n, b in model.named_buffers() if b is not None and b.is_meta] | |
| if meta_params or meta_buffers: | |
| raise RuntimeError( | |
| f'meta tensors remain after streaming load: params={meta_params[:20]} buffers={meta_buffers[:20]}' | |
| ) | |
| dense_target_weights = [ | |
| n for n, p in model.named_parameters() | |
| if n.endswith('.weight') and n in {v['official_key'] for v in manifest['targets'].values()} | |
| ] | |
| if dense_target_weights: | |
| raise RuntimeError(f'dense target weights unexpectedly materialized: {dense_target_weights[:8]}') | |
| report = { | |
| 'constructor': { | |
| 'dim': int(kwargs.get('dim', 0)), | |
| 'ffn_dim': int(kwargs.get('ffn_dim', 0)), | |
| 'num_heads': int(kwargs.get('num_heads', 0)), | |
| 'num_layers': int(kwargs.get('num_layers', 0)), | |
| 'log_scale': float(kwargs.get('log_scale', 0.0)), | |
| }, | |
| 'replaced_packed_linears': len(replaced), | |
| 'streamed_dense_non_target_tensors': assigned, | |
| 'self_contained_runtime_artifact': True, | |
| 'skipped_dense_target_weights': skipped_w4, | |
| 'meta_parameter_count': len(meta_params), | |
| 'meta_buffer_count': len(meta_buffers), | |
| 'dense_target_weight_count': len(dense_target_weights), | |
| 'a4_cache_scope': 'one_transformer_block', | |
| } | |
| model._orbitquant_w4a4_report = report | |
| model._orbitquant_engine = engine | |
| model._orbitquant_cache_clear_handles = cache_clear_handles | |
| return model, report | |
| def module_storage_bytes(module: torch.nn.Module) -> int: | |
| seen = set() | |
| total = 0 | |
| for t in list(module.parameters()) + list(module.buffers()): | |
| if t is None or t.is_meta: | |
| continue | |
| ptr = t.untyped_storage().data_ptr() | |
| if ptr in seen: | |
| continue | |
| seen.add(ptr) | |
| total += t.untyped_storage().nbytes() | |
| return total | |
| class BlockStreamer: | |
| """Exact single-GPU fallback: only one transformer block is GPU resident.""" | |
| def __init__(self, model: torch.nn.Module, device: str | torch.device = 'cuda'): | |
| self.model = model | |
| self.device = torch.device(device) | |
| self.handles = [] | |
| def install(self): | |
| # Non-block stem/head stay resident; blocks stream CPU -> GPU -> CPU. | |
| for name in ('patch_embedding', 'text_embedding', 'time_embedding', 'time_projection', 'head', 'img_emb'): | |
| mod = getattr(self.model, name, None) | |
| if mod is not None: | |
| mod.to(self.device) | |
| for block in self.model.blocks: | |
| block.to('cpu') | |
| self.handles.append(block.register_forward_pre_hook(self._pre)) | |
| self.handles.append(block.register_forward_hook(self._post)) | |
| return self | |
| def _pre(self, module, args): | |
| module.to(self.device, non_blocking=False) | |
| def _post(self, module, args, output): | |
| module.to('cpu', non_blocking=False) | |
| return output | |
| def remove(self): | |
| for h in self.handles: | |
| h.remove() | |
| self.handles.clear() | |
| def place_transformer(model, mode: str = 'auto', device: str = 'cuda', reserve_gib: float = 6.0): | |
| dev = torch.device(device) | |
| if dev.type != 'cuda': | |
| model.to(dev) | |
| return {'mode': 'resident', 'storage_bytes': module_storage_bytes(model), 'streamer': None} | |
| size = module_storage_bytes(model) | |
| free, total = torch.cuda.mem_get_info(dev) | |
| chosen = mode | |
| if mode == 'auto': | |
| chosen = 'resident' if size + int(reserve_gib * 2**30) < int(free * 0.92) else 'stream' | |
| if chosen == 'resident': | |
| model.to(dev) | |
| streamer = None | |
| elif chosen == 'stream': | |
| streamer = BlockStreamer(model, dev).install() | |
| else: | |
| raise ValueError("mode must be 'auto', 'resident', or 'stream'") | |
| return { | |
| 'mode': chosen, | |
| 'storage_bytes': size, | |
| 'free_before_bytes': int(free), | |
| 'total_vram_bytes': int(total), | |
| 'streamer': streamer, | |
| } | |