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/prefeature_hooks.py from Cccccz/HY: direct link, hf CLI and curl.
- Browser
- Download file 7.61 kB
-
https://huggingface.co/Cccccz/HY/resolve/main/predictor_data/prefeature_hooks.py
- Command line
-
hf download hf://Cccccz/HY/predictor_data/prefeature_hooks.py
-
curl -L -o prefeature_hooks.py https://huggingface.co/Cccccz/HY/resolve/main/predictor_data/prefeature_hooks.py
7.61 kB
| """Read-only Full-DiT capture of targets and joint-window Context prefeatures.""" | |
| from __future__ import annotations | |
| from typing import Any, Callable | |
| import torch | |
| from .prefeature_schema import CONTEXT_BLOCK_IDS | |
| from .v2_hooks import PredictorV2TeacherCapture, _clone_cpu | |
| class PredictorPrefeatureTeacherCapture(PredictorV2TeacherCapture): | |
| """Capture ``img_modulated`` before selected Teacher K/V projections.""" | |
| 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 = 4, | |
| context_block_ids: tuple[int, ...] = CONTEXT_BLOCK_IDS, | |
| capture_direct_kv: bool = False, | |
| ) -> None: | |
| super().__init__( | |
| transformer, | |
| on_chunk=on_chunk, | |
| num_steps=num_steps, | |
| context_block_ids=context_block_ids, | |
| ) | |
| self.capture_direct_kv = bool(capture_direct_kv) | |
| self.pending_context: dict[int, dict[str, torch.Tensor]] | None = None | |
| self.pending_direct_kv: dict[int, dict[str, torch.Tensor]] | None = None | |
| self._prefill_active = False | |
| self._prefill_features: dict[int, torch.Tensor] = {} | |
| self._prefill_metadata: dict[str, torch.Tensor] = {} | |
| def __enter__(self) -> "PredictorPrefeatureTeacherCapture": | |
| super().__enter__() | |
| for block_id in self.context_block_ids: | |
| block = self.transformer.double_blocks[block_id] | |
| self._handles.append( | |
| block.img_attn_k.register_forward_pre_hook( | |
| self._make_k_pre_hook(block_id), with_kwargs=True | |
| ) | |
| ) | |
| return self | |
| def _make_k_pre_hook(self, block_id: int): | |
| def hook(module, args, kwargs) -> None: | |
| if not self._prefill_active: | |
| return | |
| if block_id in self._prefill_features: | |
| raise RuntimeError(f"Context block {block_id} was captured twice") | |
| if not args or not torch.is_tensor(args[0]): | |
| raise RuntimeError(f"Missing img_modulated input for block {block_id}") | |
| self._prefill_features[block_id] = _clone_cpu(args[0]) | |
| return hook | |
| def _transformer_pre(self, module, args, kwargs) -> None: | |
| if kwargs.get("ar_vision_inference", False) and kwargs.get("cache_vision", False): | |
| if self.pending_context is not None or self._prefill_active: | |
| raise RuntimeError("A Context prefill was not consumed before the next prefill") | |
| frame_indices = kwargs.get("context_frame_indices") | |
| if frame_indices is None: | |
| raise RuntimeError("Context prefill lacks context_frame_indices metadata") | |
| if not torch.is_tensor(frame_indices): | |
| frame_indices = torch.tensor(frame_indices, dtype=torch.int64) | |
| frame_indices = frame_indices.detach().to(device="cpu", dtype=torch.int64).reshape(-1) | |
| frames = int(kwargs["hidden_states"].shape[2]) | |
| if frame_indices.numel() != frames: | |
| raise ValueError("Context frame metadata does not match prefill tensor") | |
| self._prefill_metadata = { | |
| "selected_frame_indices": frame_indices.contiguous(), | |
| "context_viewmats": _clone_cpu(kwargs["viewmats"]), | |
| "context_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._prefill_features = {} | |
| self._prefill_active = True | |
| return | |
| super()._transformer_pre(module, args, kwargs) | |
| if self.active is not None and self.active.step_id == 0: | |
| self.chunk_context = self._consume_context() | |
| def _transformer_post(self, module, args, kwargs, output) -> None: | |
| if self._prefill_active: | |
| self._prefill_active = False | |
| missing = set(self.context_block_ids).difference(self._prefill_features) | |
| if missing: | |
| raise RuntimeError(f"Missing prefeatures for blocks {sorted(missing)}") | |
| token_counts = {int(value.shape[1]) for value in self._prefill_features.values()} | |
| if len(token_counts) != 1: | |
| raise RuntimeError("Selected Context blocks have different window lengths") | |
| tokens = next(iter(token_counts)) | |
| mask = torch.ones((1, tokens), dtype=torch.bool) | |
| self.pending_context = { | |
| block_id: { | |
| "img_modulated": feature, | |
| "context_valid_mask": mask.clone(), | |
| **{name: value.clone() for name, value in self._prefill_metadata.items()}, | |
| } | |
| for block_id, feature in self._prefill_features.items() | |
| } | |
| if self.capture_direct_kv: | |
| self.pending_direct_kv = { | |
| block_id: { | |
| "k_vision": _clone_cpu(output[block_id]["k_vision"]), | |
| "v_vision": _clone_cpu(output[block_id]["v_vision"]), | |
| } | |
| for block_id in self.context_block_ids | |
| } | |
| self._prefill_features = {} | |
| self._prefill_metadata = {} | |
| return | |
| super()._transformer_post(module, args, kwargs, output) | |
| def _consume_context(self) -> dict[int, dict[str, torch.Tensor]]: | |
| if self.pending_context is not None: | |
| context = self.pending_context | |
| self.pending_context = None | |
| return context | |
| # The first chunk has no history prefill. | |
| empty_metadata = { | |
| "context_valid_mask": torch.ones((1, 0), dtype=torch.bool), | |
| "selected_frame_indices": torch.empty((0,), dtype=torch.int64), | |
| "context_viewmats": torch.empty((1, 0, 4, 4), dtype=torch.bfloat16), | |
| "context_Ks": torch.empty((1, 0, 3, 3), dtype=torch.bfloat16), | |
| "rope_temporal_size": torch.tensor([0], dtype=torch.int64), | |
| "start_rope_start_idx": torch.tensor([0], dtype=torch.int64), | |
| } | |
| return { | |
| block_id: { | |
| "img_modulated": torch.empty((1, 0, 2048), dtype=torch.bfloat16), | |
| **{name: value.clone() for name, value in empty_metadata.items()}, | |
| } | |
| for block_id in self.context_block_ids | |
| } | |
| def _capture_context(self, kv_cache): | |
| """Capture only case-level text KV; Context features come from block hooks.""" | |
| if self.text_context is None: | |
| text_context: dict[int, dict[str, torch.Tensor]] = {} | |
| for block_id in self.context_block_ids: | |
| cache = kv_cache[block_id] | |
| k_txt, v_txt = cache.get("k_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") | |
| text_context[block_id] = { | |
| "k_txt": _clone_cpu(k_txt), | |
| "v_txt": _clone_cpu(v_txt), | |
| } | |
| self.text_context = text_context | |
| return {block_id: {} for block_id in self.context_block_ids} | |
| def take_direct_kv(self) -> dict[int, dict[str, torch.Tensor]] | None: | |
| value = self.pending_direct_kv | |
| self.pending_direct_kv = None | |
| return value | |