Instructions to use Cccccz/HY with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use Cccccz/HY with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("Cccccz/HY", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download predictor_data/v2_hooks.py from Cccccz/HY: direct link, hf CLI and curl.
- Browser
- Download file 9.31 kB
-
https://huggingface.co/Cccccz/HY/resolve/main/predictor_data/v2_hooks.py
- Command line
-
hf download hf://Cccccz/HY/predictor_data/v2_hooks.py
-
curl -L -o v2_hooks.py https://huggingface.co/Cccccz/HY/resolve/main/predictor_data/v2_hooks.py
9.31 kB
| """Read-only capture of exact v2 trajectories and selected AR context KV.""" | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from typing import Any, Callable, Mapping | |
| import torch | |
| from .schema import LATENT_HEIGHT, LATENT_WIDTH, NUM_STEPS | |
| from .v2_schema import CONTEXT_BLOCK_IDS, block_key | |
| def _clone_cpu(tensor: torch.Tensor) -> torch.Tensor: | |
| return tensor.detach().to(device="cpu").contiguous() | |
| class _ActiveStep: | |
| chunk_id: int | |
| step_id: int | |
| tensors: dict[str, torch.Tensor] | |
| shared: dict[str, torch.Tensor] | |
| class PredictorV2TeacherCapture: | |
| """Capture Full-DiT targets plus step-invariant selected-layer AR memory.""" | |
| def __init__( | |
| self, | |
| transformer: torch.nn.Module, | |
| *, | |
| on_chunk: Callable[ | |
| [int, dict[str, torch.Tensor], dict[int, dict[str, torch.Tensor]]], None | |
| ], | |
| num_steps: int = NUM_STEPS, | |
| context_block_ids: tuple[int, ...] = CONTEXT_BLOCK_IDS, | |
| ) -> None: | |
| self.transformer = transformer | |
| self.on_chunk = on_chunk | |
| self.num_steps = num_steps | |
| self.context_block_ids = tuple(context_block_ids) | |
| self.call_index = 0 | |
| self.active: _ActiveStep | None = None | |
| self.chunk_steps: list[_ActiveStep] = [] | |
| self.image_condition_latent: torch.Tensor | None = None | |
| self.text_context: dict[int, dict[str, torch.Tensor]] | None = None | |
| self.chunk_context: dict[int, dict[str, torch.Tensor]] | None = None | |
| self._handles: list[Any] = [] | |
| def __enter__(self) -> "PredictorV2TeacherCapture": | |
| self._handles.append( | |
| self.transformer.register_forward_pre_hook(self._transformer_pre, with_kwargs=True) | |
| ) | |
| self._handles.append( | |
| self.transformer.register_forward_hook(self._transformer_post, with_kwargs=True) | |
| ) | |
| self._handles.append( | |
| self.transformer.final_layer.register_forward_pre_hook(self._final_pre, with_kwargs=True) | |
| ) | |
| return self | |
| def __exit__(self, exc_type, exc, traceback) -> bool: | |
| for handle in self._handles: | |
| handle.remove() | |
| self._handles.clear() | |
| self.active = None | |
| return False | |
| def _capture_context( | |
| self, | |
| kv_cache: list[Mapping[str, torch.Tensor | None]], | |
| ) -> dict[int, dict[str, torch.Tensor]]: | |
| result: dict[int, dict[str, torch.Tensor]] = {} | |
| for block_id in self.context_block_ids: | |
| cache = kv_cache[block_id] | |
| k_txt = cache.get("k_txt") | |
| v_txt = cache.get("v_txt") | |
| if k_txt is None or v_txt is None: | |
| raise RuntimeError(f"Block {block_id} text KV is unavailable") | |
| k_vision = cache.get("k_vision") | |
| v_vision = cache.get("v_vision") | |
| if (k_vision is None) != (v_vision is None): | |
| raise RuntimeError(f"Block {block_id} has incomplete vision KV") | |
| if k_vision is None: | |
| empty_shape = (2, int(k_txt.shape[1]), 0, int(k_txt.shape[3])) | |
| k_vision_cpu = torch.empty(empty_shape, dtype=k_txt.dtype, device="cpu") | |
| v_vision_cpu = torch.empty(empty_shape, dtype=v_txt.dtype, device="cpu") | |
| else: | |
| k_vision_cpu = _clone_cpu(k_vision) | |
| v_vision_cpu = _clone_cpu(v_vision) | |
| result[block_id] = { | |
| "k_vision": k_vision_cpu, | |
| "v_vision": v_vision_cpu, | |
| } | |
| if self.text_context is None: | |
| result[block_id]["k_txt"] = _clone_cpu(k_txt) | |
| result[block_id]["v_txt"] = _clone_cpu(v_txt) | |
| if self.text_context is None: | |
| self.text_context = { | |
| block_id: { | |
| "k_txt": result[block_id].pop("k_txt"), | |
| "v_txt": result[block_id].pop("v_txt"), | |
| } | |
| for block_id in self.context_block_ids | |
| } | |
| return result | |
| def _transformer_pre(self, module, args, kwargs) -> None: | |
| is_denoise = ( | |
| kwargs.get("ar_vision_inference", False) | |
| and not kwargs.get("cache_vision", False) | |
| ) | |
| if not is_denoise: | |
| return | |
| if self.active is not None: | |
| raise RuntimeError("Nested denoising capture is not supported") | |
| chunk_id, step_id = divmod(self.call_index, self.num_steps) | |
| model_input = kwargs["hidden_states"] | |
| if model_input.shape[1] != 65: | |
| raise ValueError(f"Teacher denoising input must have 65 channels, got {model_input.shape}") | |
| if self.image_condition_latent is None: | |
| self.image_condition_latent = _clone_cpu(model_input[:, 32:64, 0:1]) | |
| mask = model_input[:, 64:65] | |
| if not torch.all(mask[:, :, 0] == 1) or not torch.all(mask[:, :, 1:] == 0): | |
| raise ValueError("Unexpected I2V condition mask in first chunk") | |
| if step_id == 0: | |
| if self.chunk_context is not None: | |
| raise RuntimeError("Previous chunk context was not flushed") | |
| self.chunk_context = self._capture_context(kwargs["kv_cache"]) | |
| timestep = kwargs["timestep"].reshape(-1)[0:1] | |
| shared = { | |
| "action_labels": _clone_cpu(kwargs["action"].reshape(1, -1).round().long()), | |
| "target_viewmats": _clone_cpu(kwargs["viewmats"]), | |
| "target_Ks": _clone_cpu(kwargs["Ks"]), | |
| "rope_temporal_size": torch.tensor([int(kwargs["rope_temporal_size"])], dtype=torch.int64), | |
| "start_rope_start_idx": torch.tensor( | |
| [int(kwargs["start_rope_start_idx"])], dtype=torch.int64 | |
| ), | |
| } | |
| self.active = _ActiveStep( | |
| chunk_id=chunk_id, | |
| step_id=step_id, | |
| tensors={ | |
| "timestep": _clone_cpu(timestep.float()), | |
| "noisy_sample": _clone_cpu(model_input[:, :32]), | |
| }, | |
| shared=shared, | |
| ) | |
| def _final_pre(self, module, args, kwargs) -> None: | |
| if self.active is None: | |
| return | |
| hidden, condition = args[0], args[1] | |
| batch, tokens, hidden_size = hidden.shape | |
| spatial_tokens = LATENT_HEIGHT * LATENT_WIDTH | |
| if tokens % spatial_tokens: | |
| raise ValueError(f"Final hidden token count {tokens} is not frame-aligned") | |
| frames = tokens // spatial_tokens | |
| compact = condition.reshape(batch, frames, spatial_tokens, hidden_size)[:, :, 0] | |
| expanded = compact[:, :, None].expand(batch, frames, spatial_tokens, hidden_size) | |
| if not torch.equal( | |
| expanded.reshape(batch, tokens, hidden_size), | |
| condition.reshape(batch, tokens, hidden_size), | |
| ): | |
| raise ValueError("Final-layer condition varies inside a latent frame") | |
| self.active.tensors["frame_condition"] = _clone_cpu(compact) | |
| self.active.tensors["final_hidden"] = _clone_cpu(hidden) | |
| def _transformer_post(self, module, args, kwargs, output) -> None: | |
| if self.active is None: | |
| return | |
| velocity = output[0] if isinstance(output, tuple) else output | |
| self.active.tensors["velocity"] = _clone_cpu(velocity) | |
| required = {"timestep", "noisy_sample", "frame_condition", "final_hidden", "velocity"} | |
| missing = required.difference(self.active.tensors) | |
| if missing: | |
| raise RuntimeError(f"Incomplete v2 teacher step capture: {sorted(missing)}") | |
| self.chunk_steps.append(self.active) | |
| completed_step = self.active.step_id | |
| self.active = None | |
| self.call_index += 1 | |
| if completed_step == self.num_steps - 1: | |
| self._flush_chunk() | |
| def _flush_chunk(self) -> None: | |
| if len(self.chunk_steps) != self.num_steps: | |
| raise RuntimeError(f"Expected {self.num_steps} steps, got {len(self.chunk_steps)}") | |
| if self.chunk_context is None: | |
| raise RuntimeError("Chunk context was not captured") | |
| chunk_id = self.chunk_steps[0].chunk_id | |
| if any(step.chunk_id != chunk_id for step in self.chunk_steps): | |
| raise RuntimeError("Captured steps cross a chunk boundary") | |
| tensors = dict(self.chunk_steps[0].shared) | |
| for step in self.chunk_steps: | |
| for name, tensor in step.tensors.items(): | |
| tensors[f"step_{step.step_id}_{name}"] = tensor | |
| self.on_chunk(chunk_id, tensors, self.chunk_context) | |
| self.chunk_steps.clear() | |
| self.chunk_context = None | |
| def case_tensors(self) -> dict[str, torch.Tensor]: | |
| if self.image_condition_latent is None or self.text_context is None: | |
| raise RuntimeError("Case condition or text context has not been captured") | |
| result = {"image_condition_latent": self.image_condition_latent} | |
| for block_id, cache in self.text_context.items(): | |
| result[block_key(block_id, "k_txt")] = cache["k_txt"] | |
| result[block_key(block_id, "v_txt")] = cache["v_txt"] | |
| return result | |
| def text_token_count(self) -> int: | |
| if self.text_context is None: | |
| raise RuntimeError("Text context has not been captured") | |
| first = self.text_context[self.context_block_ids[0]]["k_txt"] | |
| return int(first.shape[2]) | |