diff --git a/src/h3_blackwell_runtime/packing.py b/src/h3_blackwell_runtime/packing.py index 4cd5596..511a7cd 100644 --- a/src/h3_blackwell_runtime/packing.py +++ b/src/h3_blackwell_runtime/packing.py @@ -34,6 +34,13 @@ def _axis(dim: int, area: float) -> torch.Tensor: return (torch.arange(count, dtype=torch.float64) * ratio / count + (1 - ratio) / 2) * 32 +def _frame_positions(height: int, width: int) -> torch.Tensor: + """(t ignored) area-normalized (h, w) grid of one latent frame's 2x2-patch rows.""" + area = math.sqrt(height * width) + ys, xs = torch.meshgrid(_axis(height, area), _axis(width, area), indexing="ij") + return torch.stack((ys.flatten(), xs.flatten()), dim=-1) # [frame_rows, 2] + + def _video_positions(frames: int, height: int, width: int, offset: float) -> torch.Tensor: area = math.sqrt(height * width) ys, xs = torch.meshgrid(_axis(height, area), _axis(width, area), indexing="ij") @@ -45,6 +52,18 @@ def _video_positions(frames: int, height: int, width: int, offset: float) -> tor return result.reshape(-1, 3) +def _cond_positions(frames: int, cond_t: float, height: int, width: int) -> torch.Tensor: + spatial = _frame_positions(height, width) + g = torch.empty(spatial.shape[0], 3, dtype=torch.float64) + g[:, 0] = cond_t + g[:, 1:] = spatial + return g + + +def _video_t_spans(n: int) -> list[float]: + return [FRAME_RESCALE * FRAME_PER_TOKEN[k % 5] for k in range(n)] + + def _audio_positions(steps: int, offset: float, width: int, height: int) -> torch.Tensor: area = math.sqrt(height * width) x_axis = _axis(width, area) @@ -65,6 +84,8 @@ class H3PromptPacker: self.text_weight = checkpoint.tensor("condition_proj.weight", dtype=torch.bfloat16) self.text_bias = checkpoint.tensor("condition_proj.bias", dtype=torch.bfloat16) + VISUAL_COND_TIMESTEP = 0.999 + def __call__( self, text: torch.Tensor, @@ -72,35 +93,114 @@ class H3PromptPacker: audio: torch.Tensor, sigma: float | torch.Tensor, model_timesteps: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, list[tuple[int, int, int]], tuple[int, int, int], tuple[int, int, int]]: + *, + text_token_tags: torch.Tensor | None = None, + cond_latents: list[torch.Tensor] | None = None, + frame_count: int | None = None, + seed: int = 0, + ) -> tuple[torch.Tensor, torch.Tensor, list[tuple[int, int, int]], torch.Tensor, tuple[int, int, int], tuple[int, int, int]]: + """Build ``[text | cond | audio | video]`` packed rows for (fl2va) H3. + + ``text`` is the refined text span (width 5376 when already refined, 5120 + for raw Qwen states); ``text_token_tags`` is the per-token DiT modality tag + (1=text, 0=video over vision pads). ``cond_latents`` are normalized keyframe + latents ``[1,24,1,H/16,W/16]`` spliced right after the text as non-denoised + cond rows with their own near-1 timestep. Returns + ``(hidden, times, segments, positions, video_seg, audio_seg)`` where + ``segments`` rows are ``t_row*3 + modality_tag``. + """ if text.shape[-1] == 5120: text_rows = functional.linear(text[0].to(self.text_weight.dtype), self.text_weight, self.text_bias).to(torch.bfloat16) elif text.shape[-1] == 5376: text_rows = text[0].to(torch.bfloat16) else: raise ValueError("H3 text states must be Qwen 5120-wide or refined 5376-wide.") + text_length = text_rows.shape[0] + latent_t, latent_h, latent_w = video.shape[2], video.shape[-2], video.shape[-1] + frame_rows = (latent_h // 2) * (latent_w // 2) + video_rows = functional.linear(patchify_video(video.to(torch.bfloat16)).float(), self.video_weight, self.video_bias).to(torch.bfloat16) audio_rows = functional.linear(pack_audio(audio.to(torch.bfloat16)).float(), self.audio_weight, self.audio_bias).to(torch.bfloat16) - text_length, audio_length = text_rows.shape[0], audio_rows.shape[0] - hidden = torch.cat((text_rows, audio_rows, video_rows)) + cond_rows = None + if cond_latents: + cond_patches = [] + # every cond video restarts the same CPU RNG stream (Comfy _cond_video_rows) + for idx, z in enumerate(cond_latents): + r = patchify_video(z.to(torch.float32)) + if self.VISUAL_COND_TIMESTEP < 1.0: + gen = torch.Generator("cpu").manual_seed(int(seed)) + noise = torch.randn(r.shape, generator=gen, dtype=torch.float32) + r = self.VISUAL_COND_TIMESTEP * r + (1.0 - self.VISUAL_COND_TIMESTEP) * noise.to(r.device) + cond_patches.append(r) + cond_rows = functional.linear(torch.cat(cond_patches, dim=0), self.video_weight.to(torch.float32), self.video_bias.to(torch.float32)).to(torch.bfloat16) + if model_timesteps is None: - video_sigma = torch.as_tensor(sigma, device=hidden.device, dtype=torch.float32).clamp(min=1e-6) + video_sigma = torch.as_tensor(sigma, device=text_rows.device, dtype=torch.float32).clamp(min=1e-6) base = video_sigma / (12.0 + video_sigma * (1.0 - 12.0)) audio_sigma = 3.0 * base / (1.0 + (3.0 - 1.0) * base) video_time, audio_time = (1.0 - video_sigma).item(), (1.0 - audio_sigma).item() - unique_times = sorted({video_time, audio_time}) else: - times_override = model_timesteps.to(device=hidden.device, dtype=torch.float32).flatten() - if times_override.numel() not in (1, 2): - raise ValueError("Prompt-only H3 expects one or two model timesteps.") - unique_times = times_override.tolist() - video_time, audio_time = unique_times[0], unique_times[-1] - row = {value: index for index, value in enumerate(unique_times)} - video_row, audio_row = row[video_time] * 3, row[audio_time] * 3 - times = torch.tensor(unique_times, device=hidden.device, dtype=torch.float32) - positions = torch.cat((torch.stack((torch.arange(text_length, dtype=torch.float64), torch.zeros(text_length), torch.zeros(text_length)), dim=-1), _audio_positions(audio.shape[-1], float(text_length), video.shape[-1], video.shape[-2]), _video_positions(video.shape[2], video.shape[-2], video.shape[-1], float(text_length)))) - block_video_segment = (text_length + audio_length, hidden.shape[0], video_row) - block_audio_segment = (text_length, text_length + audio_length, audio_row + 2) - final_video_segment = (text_length + audio_length, hidden.shape[0], row[video_time]) - final_audio_segment = (text_length, text_length + audio_length, row[audio_time]) - return hidden, times, [(0, text_length, video_row + 1), block_audio_segment, block_video_segment], positions, final_video_segment, final_audio_segment + times_override = model_timesteps.to(device=text_rows.device, dtype=torch.float32).flatten() + unique_override = sorted(times_override.tolist()) + video_time, audio_time = unique_override[0], unique_override[-1] + + has_vis_cond = cond_rows is not None + cond_time = max(video_time, self.VISUAL_COND_TIMESTEP) + unique_times = sorted({video_time, audio_time} | ({cond_time} if has_vis_cond else set())) + t_row = {value: index for index, value in enumerate(unique_times)} + times = torch.tensor(unique_times, device=text_rows.device, dtype=torch.float32) + + # Assemble hidden in segment order: text | cond | audio | video. + parts = [text_rows] + offsets = [0] + for block in ((cond_rows, "cond"), (audio_rows, "audio"), (video_rows, "video")): + if block[0] is None: + continue + parts.append(block[0]) + offsets.append(offsets[-1] + block[0].shape[0]) + hidden = torch.cat(parts, dim=0) + + audio_length = audio_rows.shape[0] + cond_length = cond_rows.shape[0] if cond_rows is not None else 0 + + # Positions: text rows, cond (first/last t anchors), audio, video. + text_positions = torch.stack((torch.arange(text_length, dtype=torch.float64), torch.zeros(text_length), torch.zeros(text_length)), dim=-1) + position_blocks = [text_positions] + if cond_rows is not None and cond_latents: + spans = _video_t_spans(latent_t) + cond_t_values = [ + float(text_length) if idx == 0 else (float(text_length) + sum(spans) - FRAME_RESCALE if frame_count is not None and idx == len(cond_latents) - 1 else float(text_length)) + for idx in range(len(cond_latents)) + ] + position_blocks.append(torch.cat([_cond_positions(frame_rows, cond_t, latent_h, latent_w) for cond_t in cond_t_values], dim=0)) + position_blocks.append(_audio_positions(audio.shape[-1], float(text_length + cond_length), latent_w, latent_h)) + position_blocks.append(_video_positions(latent_t, latent_h, latent_w, float(text_length + cond_length + audio_length))) + positions = torch.cat(position_blocks, dim=0) + + # mod_segments: (start, stop, t_row*3 + tag). + segments: list[tuple[int, int, int]] = [] + if text_token_tags is not None: + tags = text_token_tags.view(-1).tolist() + row_base = t_row[video_time] * 3 + run_start = 0 + for i in range(1, text_length + 1): + if i == text_length or tags[i] != tags[run_start]: + segments.append((run_start, i, row_base + int(tags[run_start]))) + run_start = i + else: + segments.append((0, text_length, t_row[video_time] * 3 + 1)) + + cursor_start = text_length + if cond_rows is not None: + segments.append((cursor_start, cursor_start + cond_rows.shape[0], t_row[cond_time] * 3 + 0)) + cursor_start += cond_rows.shape[0] + segments.append((cursor_start, cursor_start + audio_length, t_row[audio_time] * 3 + 2)) + cursor_start += audio_length + video_start = cursor_start + video_stop = video_start + video_rows.shape[0] + segments.append((video_start, video_stop, t_row[video_time] * 3 + 0)) + + video_segment = (video_start, video_stop, t_row[video_time]) + audio_start = text_length + cond_length + audio_segment = (audio_start, audio_start + audio_length, t_row[audio_time]) + return hidden, times, segments, positions, video_segment, audio_segment diff --git a/src/h3_blackwell_runtime/qwen3vl_text.py b/src/h3_blackwell_runtime/qwen3vl_text.py index 0388086..44c8b46 100644 --- a/src/h3_blackwell_runtime/qwen3vl_text.py +++ b/src/h3_blackwell_runtime/qwen3vl_text.py @@ -42,6 +42,24 @@ class _RMSNorm(nn.Module): return F.rms_norm(x, self.weight.shape, weight=self.weight.to(x), eps=self.eps) +def _apply_rope(query: torch.Tensor, key: torch.Tensor, freqs) -> tuple[torch.Tensor, torch.Tensor]: + """Direct PyTorch port of Comfy's `apply_rope` fed `freqs=(cos, sin, neg_sin)`. + + ``query``/``key`` are ``[batch, heads, seq, head_dim]``; ``freqs`` broadcast over + that shape. This is the exact split-half kernel the reference uses for both the + plain text rope and the Qwen3-VL interleaved mrope. + """ + cosine, sine, negative_sine = freqs + split = query.shape[-1] // 2 + q = query * cosine + q[..., :split].addcmul_(query[..., split:], negative_sine) + q[..., split:].addcmul_(query[..., :split], sine) + k = key * cosine + k[..., :split].addcmul_(key[..., split:], negative_sine) + k[..., split:].addcmul_(key[..., :split], sine) + return q.to(query.dtype), k.to(key.dtype) + + def _rope(query: torch.Tensor, key: torch.Tensor, theta: float) -> tuple[torch.Tensor, torch.Tensor]: """Direct PyTorch port of Comfy's text-only `precompute_freqs_cis` / `apply_rope`.""" sequence, head_dim = query.shape[-2:] @@ -54,16 +72,7 @@ def _rope(query: torch.Tensor, key: torch.Tensor, theta: float) -> tuple[torch.T sine = embedding.sin().unsqueeze(1) negative_sine = -sine[..., sine.shape[-1] // 2 :] sine = sine[..., : sine.shape[-1] // 2] - - query_output = query * cosine - split = query_output.shape[-1] // 2 - query_output[..., :split].addcmul_(query[..., split:], negative_sine) - query_output[..., split:].addcmul_(query[..., :split], sine) - key_output = key * cosine - split = key_output.shape[-1] // 2 - key_output[..., :split].addcmul_(key[..., split:], negative_sine) - key_output[..., split:].addcmul_(key[..., :split], sine) - return query_output.to(query.dtype), key_output.to(key.dtype) + return _apply_rope(query, key, (cosine, sine, negative_sine)) def _qwen_attention(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: @@ -105,7 +114,7 @@ class _Qwen3VLBlock(nn.Module): self.up_proj = checkpoint.nvfp4_linear(f"{prefix}.mlp.up_proj", output_dtype=dtype) self.down_proj = checkpoint.nvfp4_linear(f"{prefix}.mlp.down_proj", output_dtype=dtype) - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + def forward(self, hidden_states: torch.Tensor, freqs_cis=None) -> torch.Tensor: residual = hidden_states x = self.input_layernorm(hidden_states) batch, sequence, _ = x.shape @@ -114,7 +123,10 @@ class _Qwen3VLBlock(nn.Module): value = self.v_proj(x).view(batch, sequence, self.config.num_key_value_heads, self.config.head_dim).transpose(1, 2) query = self.q_norm(query) key = self.k_norm(key) - query, key = _rope(query, key, self.config.rope_theta) + if freqs_cis is None: + query, key = _rope(query, key, self.config.rope_theta) + else: + query, key = _apply_rope(query, key, freqs_cis) # Comfy selects its small-input SDPA path for Qwen, with an explicit causal mask. causal_mask = torch.full( (sequence, sequence), @@ -178,6 +190,44 @@ class Qwen3VL32BTextEncoder(nn.Module): + ", ".join(missing) ) + def _embed_rows(self, input_ids: torch.Tensor) -> torch.Tensor: + """Scaled token embeds ``[batch, tokens, 5120]`` in fp32 (pre-decoder).""" + token_rows = F.embedding(input_ids.to(self.embed_tokens.device), self.embed_tokens).to(torch.float32) + token_scales = F.embedding(input_ids.to(self.embed_scale.device), self.embed_scale) + return (token_rows * token_scales).to(torch.float32) + + def _run_layers( + self, + hidden_states: torch.Tensor, + *, + position_ids: torch.Tensor | None = None, + visual_pos_masks: torch.Tensor | None = None, + deepstack_embeds: list[torch.Tensor] | None = None, + ) -> torch.Tensor: + """Run all 50 decoder blocks over an fp32 ``[batch, tokens, 5120]`` state. + + With ``position_ids`` (``[3, seq]``) the Qwen3-VL interleaved mrope is used + instead of the plain 1D text rope; with ``visual_pos_masks`` + + ``deepstack_embeds`` the three DeepStack features are added at the visual + positions of the first three decoder layers (Comfy ``Llama2_.forward``). + """ + hidden_states = hidden_states.to(self.dtype) + freqs_cis = None + if position_ids is not None: + from .qwen3vl_vision import mrope_freqs_cis + + freqs_cis = mrope_freqs_cis(position_ids, device=hidden_states.device) + for layer_index, layer in enumerate(self.layers): + hidden_states = layer(hidden_states, freqs_cis) + if ( + deepstack_embeds is not None + and visual_pos_masks is not None + and layer_index < len(deepstack_embeds) + ): + mask = visual_pos_masks.to(hidden_states.device) + hidden_states[mask] = hidden_states[mask] + deepstack_embeds[layer_index].to(hidden_states) + return hidden_states + @torch.inference_mode() def forward(self, input_ids: torch.Tensor) -> torch.Tensor: """Return unnormalized `[batch, tokens, 5120]` output after decoder layer 50.""" @@ -185,12 +235,7 @@ class Qwen3VL32BTextEncoder(nn.Module): raise ValueError(f"input_ids must have shape [batch, tokens], got {tuple(input_ids.shape)}") if input_ids.numel() == 0: raise ValueError("input_ids must contain at least one token") - token_rows = F.embedding(input_ids.to(self.embed_tokens.device), self.embed_tokens).to(torch.float32) - token_scales = F.embedding(input_ids.to(self.embed_scale.device), self.embed_scale) - hidden_states = (token_rows * token_scales).to(self.dtype) - for layer in self.layers: - hidden_states = layer(hidden_states) - return hidden_states + return self._run_layers(self._embed_rows(input_ids)) class Qwen3VLPromptConditioner: diff --git a/src/h3_blackwell_runtime/qwen3vl_vision.py b/src/h3_blackwell_runtime/qwen3vl_vision.py new file mode 100644 index 0000000..90eebb5 --- /dev/null +++ b/src/h3_blackwell_runtime/qwen3vl_vision.py @@ -0,0 +1,692 @@ +"""Qwen3-VL vision conditioning for the direct H3 runtime (first/last keyframes). + +Standalone port of the Comfy reference's Qwen3.5/Qwen3-VL visual stack and the +MiniMax H3 presentation mechanism, built only on ``torch``: + +- ``Qwen35VisionModel`` (visual.pos_embed, visual.patch_embed, 27 visual.blocks + with 2D rotary attention, visual.merger) plus the three + ``Qwen3VLDeepstackMerger`` (visual.deepstack_merger_list). +- ``process_qwen2vl_images`` image preprocessing (Qwen grid/resize/normalize). +- ``qwen2vl_mrope_position_ids`` (H3 mrope ids from ``embeds_info``), + ``precompute_mrope_freqs_cis`` (interleaved Qwen3-VL text mrope), and + ``token_tags_from_embeds_info`` (DiT per-token modality tags). +- Comfy ``Llama2_.forward`` DeepStack plumbing: per-decoder-layer additions at + visual positions for the first three decoder layers. + +The visual tower is loaded from the same ``qwen3vl_32b_minimax_h3_nvfp4_awq`` +safetensors the text encoder uses; its keys are all plain (bf16) tensors under +the ``visual.*`` prefix (verified at 1.1 GB, no quantized sub-tensors). +""" + +from __future__ import annotations + +import math +from dataclasses import dataclass +from pathlib import Path + +import torch +from safetensors import safe_open +from torch import nn +from torch.nn import functional as F + + +# H3 presentation sentinels (see upstream_qwen3vl.py / upstream_text.py). +VISION_START = 151652 +VISION_END = 151653 +IMAGE_EMBED_TOKEN = 151655 + +# Qwen3-VL-32B visual geometry (QWEN3VL_VISION in upstream_qwen3vl.py). +VISION_HIDDEN = 1152 +VISION_INTERMEDIATE = 4304 +VISION_DEPTH = 27 +VISION_HEADS = 16 +VISION_HEAD_DIM = VISION_HIDDEN // VISION_HEADS # 72 +VISION_PATCH = 16 +VISION_TEMPORAL = 2 +VISION_MERGE = 2 +VISION_POSITION_EMBEDS = 2304 +DEEPSTACK_VISUAL_INDEXES = (8, 16, 24) + +# Qwen3-VL text mrope geometry (Qwen3VL_32BConfig in llama.py). +TEXT_ROPE_DIMS = (24, 20, 20) +TEXT_ROPE_THETA = 5_000_000.0 +TEXT_HEAD_DIM = 128 + +# Qwen image preprocessing policy (process_qwen2vl_images, H3 mean/std 0.5). +QWEN_IMAGE_MEAN = (0.5, 0.5, 0.5) +QWEN_IMAGE_STD = (0.5, 0.5, 0.5) +QWEN_MIN_PIXELS = 3136 +QWEN_MAX_PIXELS = 12845056 + + +def process_image( + image: torch.Tensor, + *, + min_pixels: int = QWEN_MIN_PIXELS, + max_pixels: int = QWEN_MAX_PIXELS, + patch_size: int = VISION_PATCH, + temporal_patch_size: int = VISION_TEMPORAL, + merge_size: int = VISION_MERGE, + image_mean: tuple[float, ...] = QWEN_IMAGE_MEAN, + image_std: tuple[float, ...] = QWEN_IMAGE_STD, +) -> tuple[torch.Tensor, torch.Tensor]: + """Qwen image preprocessing (process_qwen2vl_images): a single + ``[1, H, W, 3]`` float image (``[0, 1]``) -> (flatten_patches + ``[grid_h*grid_w, C*tp*ps*ps]``, image_grid_thw ``[1, grid_h, grid_w]``). + """ + if image.ndim != 4 or image.shape[0] != 1: + raise ValueError("process_image expects a single [1, H, W, 3] tensor") + batch, height, width, _ = image.shape + device = image.device + + images = image.permute(0, 3, 1, 2) + img = images[0] + + factor = patch_size * merge_size + h_bar = round(height / factor) * factor + w_bar = round(width / factor) * factor + + if h_bar * w_bar > max_pixels: + beta = math.sqrt((height * width) / max_pixels) + h_bar = max(factor, math.floor(height / beta / factor) * factor) + w_bar = max(factor, math.floor(width / beta / factor) * factor) + elif h_bar * w_bar < min_pixels: + beta = math.sqrt(min_pixels / (height * width)) + h_bar = math.ceil(height * beta / factor) * factor + w_bar = math.ceil(width * beta / factor) * factor + + img_resized = F.interpolate(img.unsqueeze(0), size=(h_bar, w_bar), mode="bilinear", align_corners=False).squeeze(0) + + normalized = img_resized.clone() + for c in range(3): + normalized[c] = (img_resized[c] - image_mean[c]) / image_std[c] + + grid_h = h_bar // patch_size + grid_w = w_bar // patch_size + grid_thw = torch.tensor([1, grid_h, grid_w], device=device, dtype=torch.long) + + pixel_values = normalized + channel = pixel_values.shape[0] + grid_t = 1 + pixel_values = pixel_values.unsqueeze(0).repeat(2, 1, 1, 1) + + patches = pixel_values.reshape( + grid_t, + temporal_patch_size, + channel, + grid_h // merge_size, + merge_size, + patch_size, + grid_w // merge_size, + merge_size, + patch_size, + ) + patches = patches.permute(0, 3, 6, 4, 7, 2, 1, 5, 8) + flatten = patches.reshape(grid_t * grid_h * grid_w, channel * temporal_patch_size * patch_size * patch_size) + return flatten, grid_thw + + +def mrope_position_ids(embeds_info: list[dict], seq_len: int, device) -> torch.Tensor | None: + """(T, H, W) mrope ids for a sequence with spliced visual blocks (reference + ``qwen2vl_mrope_position_ids``). ``embeds_info`` entries carry + ``index``/``size`` spans and an ``extra`` dict with a ``grid`` tensor + ``[1, grid_h, grid_w]``. Returns ``[3, seq_len]`` or ``None``.""" + position_ids = None + offset = 0 + for e in embeds_info: + if e.get("type") != "image": + continue + extra = e.get("extra", None) + grid = extra["grid"] if isinstance(extra, dict) else extra + start = e.get("index") + if position_ids is None: + position_ids = torch.zeros((3, seq_len), device=device) + position_ids[:, :start] = torch.arange(0, start, device=device) + end = e.get("size") + start + len_max = int(grid.max()) // 2 + start_next = len_max + start + position_ids[:, end:] = torch.arange(start_next + offset, start_next + (seq_len - end) + offset, device=device) + position_ids[0, start:end] = start + offset + max_d = int(grid[0][1]) // 2 + position_ids[1, start:end] = ( + torch.arange(start + offset, start + max_d + offset, device=device) + .unsqueeze(1) + .repeat(1, math.ceil((end - start) / max_d)) + .flatten(0)[: end - start] + ) + max_d = int(grid[0][2]) // 2 + position_ids[2, start:end] = ( + torch.arange(start + offset, start + max_d + offset, device=device) + .unsqueeze(0) + .repeat(math.ceil((end - start) / max_d), 1) + .flatten(0)[: end - start] + ) + offset += len_max - (end - start) + return position_ids + + +def token_tags(seq_len: int, embeds_info: list[dict], device) -> torch.Tensor: + """DiT per-token AdaLN tags: 0 (video) inside a visual block including the + flanking sentinels, 1 (text) elsewhere (reference + ``token_tags_from_embeds_info``).""" + tags = torch.ones(seq_len, dtype=torch.long, device=device) + for e in embeds_info: + if e.get("type") == "image": + start = max(0, e["index"] - 1) + stop = e["index"] + e["size"] + 1 + tags[start:stop] = 0 + return tags + + +def mrope_freqs_cis(position_ids: torch.Tensor, *, theta: float = TEXT_ROPE_THETA, head_dim: int = TEXT_HEAD_DIM, rope_dims: tuple[int, ...] = TEXT_ROPE_DIMS) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Interleaved Qwen3-VL text mrope (rope_dims=(24,20,20)). + + ``position_ids`` is ``[3, seq]``; returns ``(cos, sin, neg_sin)`` matching + the reference ``precompute_freqs_cis`` + ``apply_rope`` convention (shape + ``[1, seq, head_dim]``; the ``neg_sin`` entry is ``-sin[..., half:]``). + """ + if position_ids.shape[0] < 1: + position_ids = position_ids.unsqueeze(0) + inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=position_ids.device, dtype=torch.float32) / head_dim)) + inv_freq_expanded = inv_freq[None, :, None].expand(position_ids.shape[0], -1, 1) + position_ids_expanded = position_ids[:, None, :].to(torch.float32) + freqs = (inv_freq_expanded @ position_ids_expanded).transpose(1, 2) + freqs_inter = freqs[0].clone() + for axis_idx, offset in ((1, 1), (2, 2)): + length = rope_dims[axis_idx - 1] * 3 + idx = slice(offset, length, 3) + freqs_inter[..., idx] = freqs[axis_idx, ..., idx] + emb = torch.cat((freqs_inter, freqs_inter), dim=-1) + cos = emb.cos().unsqueeze(0) + sin = emb.sin().unsqueeze(0) + sine = sin[..., : sin.shape[-1] // 2] + negative_sine = -sin[..., sin.shape[-1] // 2 :] + return cos, sine, negative_sine + + +class _VisionPatchEmbed(nn.Module): + def __init__(self, weight: torch.Tensor, bias: torch.Tensor): + super().__init__() + self.register_buffer("weight", weight, persistent=False) + self.register_buffer("bias", bias, persistent=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + target = self.weight.dtype + x = x.view(-1, 3, VISION_TEMPORAL, VISION_PATCH, VISION_PATCH) + return F.conv3d(x.to(target), self.weight, self.bias, kernel_size=(VISION_TEMPORAL, VISION_PATCH, VISION_PATCH), stride=(VISION_TEMPORAL, VISION_PATCH, VISION_PATCH)).view(-1, self.weight.shape[0]) + + +class _VisionMLP(nn.Module): + def __init__(self, fc1_w: torch.Tensor, fc1_b: torch.Tensor, fc2_w: torch.Tensor, fc2_b: torch.Tensor): + super().__init__() + self.register_buffer("fc1_weight", fc1_w, persistent=False) + self.register_buffer("fc1_bias", fc1_b, persistent=False) + self.register_buffer("fc2_weight", fc2_w, persistent=False) + self.register_buffer("fc2_bias", fc2_b, persistent=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return F.linear(F.gelu(F.linear(x, self.fc1_weight, self.fc1_bias), approximate="tanh"), self.fc2_weight, self.fc2_bias) + + +def _rotate_half(x: torch.Tensor) -> torch.Tensor: + x1 = x[..., : x.shape[-1] // 2] + x2 = x[..., x.shape[-1] // 2 :] + return torch.cat((-x2, x1), dim=-1) + + +def _apply_rope_vision(q: torch.Tensor, k: torch.Tensor, freqs) -> tuple[torch.Tensor, torch.Tensor]: + """Reference ``apply_rope`` (split-half) fed ``freqs=(cos, sin, neg_sin)``. + + The vision rotary is built from the doubled-angle embedding so that + ``sin[..., half:]`` already equals ``-sin[..., :half]``; the caller passes it + through unchanged, making this the rotate-half equivalent and keeping q/k at + their native dtype. + """ + cos, sin, neg_sin = freqs + q = (q * cos) + split = q.shape[-1] // 2 + q[..., :split] += q[..., split:] * neg_sin + q[..., split:] += q[..., :split] * sin + k = (k * cos) + k[..., :split] += k[..., split:] * neg_sin + k[..., split:] += k[..., :split] * sin + return q, k + + +class _VisionAttention(nn.Module): + def __init__(self, qkv_w: torch.Tensor, qkv_b: torch.Tensor, proj_w: torch.Tensor, *, num_heads: int, head_dim: int): + super().__init__() + self.num_heads = num_heads + self.head_dim = head_dim + self.register_buffer("qkv_weight", qkv_w, persistent=False) + self.register_buffer("qkv_bias", qkv_b, persistent=False) + self.proj_weight = proj_w # no bias (Qwen3.5 vision proj has none) + + def forward(self, x: torch.Tensor, cu_seqlens: torch.Tensor, position_embeddings: torch.Tensor) -> torch.Tensor: + seq_length = x.shape[0] + query_states, key_states, value_states = ( + F.linear(x, self.qkv_weight, self.qkv_bias).reshape(seq_length, 3, self.num_heads, self.head_dim).permute(1, 0, 2, 3).unbind(0) + ) + query_states, key_states = _apply_rope_vision(query_states, key_states, position_embeddings) + lengths = (cu_seqlens[1:] - cu_seqlens[:-1]).tolist() + attn_outputs = [] + for q, k, v in zip( + torch.split(query_states, lengths, dim=0), + torch.split(key_states, lengths, dim=0), + torch.split(value_states, lengths, dim=0), + ): + attn_outputs.append(F.scaled_dot_product_attention(q.transpose(0, 1).unsqueeze(0), k.transpose(0, 1).unsqueeze(0), v.transpose(0, 1).unsqueeze(0))) + attn_output = torch.cat(attn_outputs, dim=1) + attn_output = attn_output.reshape(seq_length, -1) + return F.linear(attn_output, self.proj_weight) + + +class _VisionBlock(nn.Module): + def __init__(self, norm1_w: torch.Tensor, norm1_b: torch.Tensor, attn: _VisionAttention, norm2_w: torch.Tensor, norm2_b: torch.Tensor, mlp: _VisionMLP): + super().__init__() + self.attn = attn + self.mlp = mlp + self.register_buffer("norm1_weight", norm1_w, persistent=False) + self.register_buffer("norm1_bias", norm1_b, persistent=False) + self.register_buffer("norm2_weight", norm2_w, persistent=False) + self.register_buffer("norm2_bias", norm2_b, persistent=False) + + def forward(self, x: torch.Tensor, cu_seqlens: torch.Tensor, position_embeddings: torch.Tensor) -> torch.Tensor: + x = x + self.attn(F.layer_norm(x, (x.shape[-1],), weight=self.norm1_weight, bias=self.norm1_bias, eps=1e-6), cu_seqlens=cu_seqlens, position_embeddings=position_embeddings) + return x + self.mlp(F.layer_norm(x, (x.shape[-1],), weight=self.norm2_weight, bias=self.norm2_bias, eps=1e-6)) + + +class _VisionPatchMerger(nn.Module): + """Spatial 2x2 merge + projection to text width (main or deepstack).""" + + def __init__(self, norm_w: torch.Tensor, norm_b: torch.Tensor, fc1_w: torch.Tensor, fc1_b: torch.Tensor, fc2_w: torch.Tensor, fc2_b: torch.Tensor, *, merge_size: int, out_hidden_size: int): + super().__init__() + self.merge_dim = VISION_HIDDEN * (merge_size ** 2) + self.register_buffer("norm_weight", norm_w, persistent=False) + self.register_buffer("norm_bias", norm_b, persistent=False) + self.register_buffer("fc1_weight", fc1_w, persistent=False) + self.register_buffer("fc1_bias", fc1_b, persistent=False) + self.register_buffer("fc2_weight", fc2_w, persistent=False) + self.register_buffer("fc2_bias", fc2_b, persistent=False) + self.out_hidden_size = out_hidden_size + + def forward(self, x: torch.Tensor) -> torch.Tensor: + # x: [seq, hidden] (already spatially merged upstream). + x = F.layer_norm(x, (VISION_HIDDEN,), weight=self.norm_weight, bias=self.norm_bias, eps=1e-6) + x = x.view(-1, self.merge_dim) + return F.linear(F.gelu(F.linear(x, self.fc1_weight, self.fc1_bias), approximate="tanh"), self.fc2_weight, self.fc2_bias) + + +def resize_keyframe(image: torch.Tensor, width: int, height: int, *, crop: str = "disabled") -> torch.Tensor: + """Resize a ``[1, H, W, 3]`` float ``[0, 1]`` keyframe to the ``width x height`` + canvas the way the Comfy reference does (``common_upscale(..., "lanczos", crop)``): + + - ``crop="disabled"`` (first frame) stretches to the canvas. + - ``crop="center"`` (last frame) aspect-covers then center-crops. + + Returns ``[1, height, width, 3]`` float ``[0, 1]``. + """ + if image.ndim != 4 or image.shape[0] != 1: + raise ValueError("resize_keyframe expects a single [1, H, W, 3] image") + samples = image[:, :, :, :3].movedim(-1, 1) # [1, 3, H, W] + if crop == "center": + old_h, old_w = samples.shape[-2], samples.shape[-1] + old_aspect = old_w / old_h + new_aspect = width / height + x = 0 + y = 0 + if old_aspect > new_aspect: + x = round((old_w - old_w * (new_aspect / old_aspect)) / 2) + elif old_aspect < new_aspect: + y = round((old_h - old_h * (old_aspect / new_aspect)) / 2) + samples = samples.narrow(-2, y, old_h - 2 * y).narrow(-1, x, old_w - 2 * x) + samples = F.interpolate(samples, size=(height, width), mode="lanczos") + return samples.clamp(0.0, 1.0).movedim(1, -1) # [1, H, W, 3] + + +def _text_run_ids(prompt: str) -> list[int]: + """Token ids for a raw text run (``add_special_tokens=False``, no template).""" + from .conditioning import H3PromptTokenizer + + tokenizer_dir = Path(__file__).with_name("qwen25_tokenizer") + if not tokenizer_dir.exists(): + raise FileNotFoundError(f"Qwen tokenizer directory missing: {tokenizer_dir}") + return H3PromptTokenizer(tokenizer_dir)(prompt or " ") + + +@dataclass +class Fl2vaPresentation: + """Expanded first/last keyframe prompt for the H3 DiT. + + Carries the Qwen layer-50 text states (with spliced vision pads), the per-token + DiT modality tags, the mrope ids, and the DiT-level keyframe anchors used to + build the cond rows. + """ + + input_ids: torch.Tensor + text_states: torch.Tensor + text_token_tags: torch.Tensor + embeds_info: list[dict] + keyframes: list[dict] + frame_count: int + + +def build_fl2va_presentation( + prompt: str, + first_frame: torch.Tensor | None, + last_frame: torch.Tensor | None, + *, + width: int, + height: int, + frame_count: int, + tokenizer, + vision: Qwen3VL32BVision, + text_encoder: "Qwen3VL32BTextEncoder", + device, +) -> Fl2vaPresentation: + """Build the fl2va presentation and run the Qwen text conditioner over it. + + ``first_frame`` / ``last_frame`` are ``[1, H, W, 3]`` ``[0, 1]`` images (already + aligned to the target canvas by the caller when needed). The presentation + sequence is ``[Picture1 Picture2? prompt]``: raw text ids with the + ``VISION_START`` sentinel, a vision pad run, and ``VISION_END`` per keyframe, + followed by the raw prompt ids (no chat template). + """ + images: list[torch.Tensor] = [] + keyframes: list[dict] = [] + if first_frame is not None: + img = resize_keyframe(first_frame, width, height, crop="disabled") + images.append(img) + keyframes.append({"resolved_frame_index": 0, "image": img}) + if last_frame is not None: + img = resize_keyframe(last_frame, width, height, crop="center") + images.append(img) + keyframes.append({"resolved_frame_index": frame_count - 1, "image": img}) + + # Build the entry list: (token_id/"text") runs and image placeholders. + entries: list = [] + for i in range(len(images)): + entries.extend((tid, "text") for tid in _text_run_ids(f" None: + self.register_buffer("pos_embed", get("visual.pos_embed.weight"), persistent=False) + self.patch_embed = _VisionPatchEmbed(get("visual.patch_embed.proj.weight"), get("visual.patch_embed.proj.bias")) + + self.merger = _VisionPatchMerger( + get("visual.merger.norm.weight"), get("visual.merger.norm.bias"), + get("visual.merger.linear_fc1.weight"), get("visual.merger.linear_fc1.bias"), + get("visual.merger.linear_fc2.weight"), get("visual.merger.linear_fc2.bias"), + merge_size=self.spatial_merge_size, out_hidden_size=self.out_hidden_size, + ) + self.deepstack_merger_list = nn.ModuleList([ + _VisionPatchMerger( + get(f"visual.deepstack_merger_list.{i}.norm.weight"), get(f"visual.deepstack_merger_list.{i}.norm.bias"), + get(f"visual.deepstack_merger_list.{i}.linear_fc1.weight"), get(f"visual.deepstack_merger_list.{i}.linear_fc1.bias"), + get(f"visual.deepstack_merger_list.{i}.linear_fc2.weight"), get(f"visual.deepstack_merger_list.{i}.linear_fc2.bias"), + merge_size=self.spatial_merge_size, out_hidden_size=self.out_hidden_size, + ) + for i in range(len(self.deepstack_visual_indexes)) + ]) + self.blocks = nn.ModuleList([ + _VisionBlock( + get(f"visual.blocks.{i}.norm1.weight"), get(f"visual.blocks.{i}.norm1.bias"), + _VisionAttention( + get(f"visual.blocks.{i}.attn.qkv.weight"), get(f"visual.blocks.{i}.attn.qkv.bias"), + get(f"visual.blocks.{i}.attn.proj.weight"), + num_heads=self.num_heads, head_dim=VISION_HEAD_DIM, + ), + get(f"visual.blocks.{i}.norm2.weight"), get(f"visual.blocks.{i}.norm2.bias"), + _VisionMLP( + get(f"visual.blocks.{i}.mlp.linear_fc1.weight"), get(f"visual.blocks.{i}.mlp.linear_fc1.bias"), + get(f"visual.blocks.{i}.mlp.linear_fc2.weight"), get(f"visual.blocks.{i}.mlp.linear_fc2.bias"), + ), + ) + for i in range(self.depth) + ]) + self.rotary_pos_emb = _VisionRotary(VISION_HIDDEN // self.num_heads // 2, device, dtype) + + def rot_pos_emb(self, grid_thw: torch.Tensor) -> torch.Tensor: + """Reference ``rot_pos_emb``: (row, col) coords -> 2D rotary angles per token.""" + merge_size = self.spatial_merge_size + grid_thw_list = grid_thw.tolist() + max_hw = max(max(h, w) for _, h, w in grid_thw_list) + freq_table = self.rotary_pos_emb(max_hw).to(grid_thw.device) + device = freq_table.device + total_tokens = sum(int(t * h * w) for t, h, w in grid_thw_list) + pos_ids = torch.empty((total_tokens, 2), dtype=torch.long, device=device) + offset = 0 + for num_frames, height, width in grid_thw_list: + num_frames, height, width = int(num_frames), int(height), int(width) + merged_h, merged_w = height // merge_size, width // merge_size + block_rows = torch.arange(merged_h, device=device) + block_cols = torch.arange(merged_w, device=device) + intra_row = torch.arange(merge_size, device=device) + intra_col = torch.arange(merge_size, device=device) + row_idx = (block_rows[:, None, None, None] * merge_size + intra_row[None, None, :, None]).expand(merged_h, merged_w, merge_size, merge_size).reshape(-1) + col_idx = (block_cols[None, :, None, None] * merge_size + intra_col[None, None, :, None].expand(merged_h, merged_w, merge_size, merge_size).reshape(-1)) + coords = torch.stack((row_idx, col_idx), dim=-1) + if num_frames > 1: + coords = coords.repeat(num_frames, 1) + num_tokens = coords.shape[0] + pos_ids[offset:offset + num_tokens] = coords + offset += num_tokens + return freq_table[pos_ids].flatten(1) + + def fast_pos_embed_interpolate(self, grid_thw: torch.Tensor) -> torch.Tensor: + """Reference 4-tap bilinear interpolation of the learned 48x48 grid.""" + grid_ts = [int(row[0]) for row in grid_thw.tolist()] + grid_hs = [int(row[1]) for row in grid_thw.tolist()] + grid_ws = [int(row[2]) for row in grid_thw.tolist()] + device = self.pos_embed.device + idx_list: list[list] = [[] for _ in range(4)] + weight_list: list[list] = [[] for _ in range(4)] + for t, h, w in zip(grid_ts, grid_hs, grid_ws): + h, w = int(h), int(w) + h_idxs = torch.linspace(0, self.num_grid_per_side - 1, h, device=device) + w_idxs = torch.linspace(0, self.num_grid_per_side - 1, w, device=device) + h_floor = h_idxs.int() + w_floor = w_idxs.int() + h_ceil = (h_idxs + 1).int().clamp(max=self.num_grid_per_side - 1) + w_ceil = (w_idxs + 1).int().clamp(max=self.num_grid_per_side - 1) + dh = h_idxs - h_floor + dw = w_idxs - w_floor + base_h = h_floor * self.num_grid_per_side + base_h_ceil = h_ceil * self.num_grid_per_side + indices = [ + (base_h[None].T + w_floor[None]).flatten(), + (base_h[None].T + w_ceil[None]).flatten(), + (base_h_ceil[None].T + w_floor[None]).flatten(), + (base_h_ceil[None].T + w_ceil[None]).flatten(), + ] + weights = [ + ((1 - dh)[None].T * (1 - dw)[None]).flatten(), + ((1 - dh)[None].T * dw[None]).flatten(), + (dh[None].T * (1 - dw)[None]).flatten(), + (dh[None].T * dw[None]).flatten(), + ] + for j in range(4): + idx_list[j].extend(indices[j].tolist()) + weight_list[j].extend(weights[j].tolist()) + idx_tensor = torch.tensor(idx_list, dtype=torch.long, device=device) + weight_tensor = torch.tensor(weight_list, dtype=self.pos_embed.dtype, device=device) + pos_embeds = self.pos_embed[idx_tensor] * weight_tensor[:, :, None] + patch_pos_embeds = pos_embeds[0] + pos_embeds[1] + pos_embeds[2] + pos_embeds[3] + patch_pos_embeds = patch_pos_embeds.split([h * w for h, w in zip(grid_hs, grid_ws)]) + patch_pos_embeds_permute = [] + merge_size = self.spatial_merge_size + for pos_embed, t, h, w in zip(patch_pos_embeds, grid_ts, grid_hs, grid_ws): + pos_embed = pos_embed.repeat(t, 1) + pos_embed = pos_embed.view(t, h // merge_size, merge_size, w // merge_size, merge_size, -1).permute(0, 1, 3, 2, 4, 5).flatten(0, 4) + patch_pos_embeds_permute.append(pos_embed) + return torch.cat(patch_pos_embeds_permute) + + @staticmethod + def _merge_tokens(x: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor: + """Flatten 2x2 spatial-merge blocks (the reference's pre-merge flatten).""" + t, h, w = (int(v) for v in grid_thw.tolist()[0]) + merge = 2 + return x.view(t, h // merge, merge, w // merge, merge, -1).permute(0, 1, 3, 2, 4, 5).reshape(-1, -1) + + def forward(self, flatten_patches: torch.Tensor, grid_thw: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]: + """Run the visual tower -> (merged, deepstack).""" + x = self.patch_embed(flatten_patches.to(self.dtype).to(self.device)) + x = x + self.fast_pos_embed_interpolate(grid_thw).to(x.device) + x = x.reshape(x.shape[0], -1) + rotary = self.rot_pos_emb(grid_thw.to(x.device)).to(x.device).reshape(x.shape[0], -1) + emb = torch.cat((rotary, rotary), dim=-1) + cos = emb.cos().unsqueeze(-2) + sin = emb.sin().unsqueeze(-2) + sin_split = sin.shape[-1] // 2 + position_embeddings = (cos, sin[..., :sin_split], -sin[..., sin_split:]) + cu_seqlens = F.pad( + torch.repeat_interleave(grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0]).cumsum(0, dtype=torch.int32), + (1, 0), value=0, + ) + deepstack_features = [] + for layer_num, block in enumerate(self.blocks): + x = block(x, cu_seqlens=cu_seqlens, position_embeddings=position_embeddings) + if layer_num in self.deepstack_visual_indexes: + deepstack_features.append(self.deepstack_merger_list[self.deepstack_visual_indexes.index(layer_num)](self._merge_tokens(x, grid_thw))) + return self.merger(self._merge_tokens(x, grid_thw)), deepstack_features + + +class _VisionRotary(nn.Module): + def __init__(self, dim: int, device, dtype): + super().__init__() + inv_freq = 1.0 / (10000.0 ** (torch.arange(0, dim, 2, dtype=torch.float, device=device) / dim)) + self.register_buffer("inv_freq", inv_freq, persistent=False) + + def forward(self, seqlen: int) -> torch.Tensor: + seq = torch.arange(seqlen, device=self.inv_freq.device, dtype=self.inv_freq.dtype) + return torch.outer(seq, self.inv_freq) diff --git a/src/h3_blackwell_runtime/runtime.py b/src/h3_blackwell_runtime/runtime.py index c14ac93..6a840f9 100644 --- a/src/h3_blackwell_runtime/runtime.py +++ b/src/h3_blackwell_runtime/runtime.py @@ -17,10 +17,12 @@ from .checkpoint import H3Checkpoint from .denoiser import H3PackedDenoiser from .packing import H3PromptPacker from .qwen3vl_text import Qwen3VLPromptConditioner +from .qwen3vl_vision import build_fl2va_presentation, Qwen3VL32BVision from .sampler import sample_video_res_multistep from .t2v import random_av_latents from .token_refiner import H3TokenRefiner from .vae_decoder import MiniMaxH3VideoVAE, dtype_from_name +from .vae_encoder import MiniMaxH3VideoVAEEncoder @dataclass(frozen=True) @@ -85,6 +87,14 @@ class H3HotRuntime: "audio_vae_loaded", lambda: MiniMaxH3AudioVAE.from_safetensors(config.audio_vae_path, device=config.device).eval(), ) + self.vae_encoder = self._timed_load( + "vae_encoder_loaded", + lambda: MiniMaxH3VideoVAEEncoder.from_safetensors(config.video_vae_path, device=config.device).eval(), + ) + self.vision_tower = self._timed_load( + "vision_tower_loaded", + lambda: Qwen3VL32BVision(config.text_encoder_path, device=config.device, dtype=torch.bfloat16), + ) def _timed_load(self, stage: str, fn): _sync() @@ -134,6 +144,44 @@ class H3HotRuntime: block.attention_backend = _refiner_attention_backend(attention) self.attention = attention + def _build_fl2va(self, prompt: str, first_frame: torch.Tensor | None, last_frame: torch.Tensor | None, width: int, height: int, aligned_frames: int): + from .qwen3vl_vision import build_fl2va_presentation, resize_keyframe + + device = self.config.device + first = self._image_to_uint8_nhwc(first_frame.to(device)).float() / 255.0 if first_frame is not None else None + last = self._image_to_uint8_nhwc(last_frame.to(device)).float() / 255.0 if last_frame is not None else None + presentation = build_fl2va_presentation( + prompt, + first, + last, + width=width, + height=height, + frame_count=aligned_frames, + tokenizer=self.conditioner.tokenizer, + vision=self.vision_tower, + text_encoder=self.conditioner.encoder, + device=device, + ) + # resize each keyframe to the canvas and VAE-encode it (the DiT cond latent). + cond_latents = [] + for kf in presentation.keyframes: + resized = resize_keyframe(kf["image"].to(device), width, height, crop="disabled" if kf["resolved_frame_index"] == 0 else "center") + pix = resized.movedim(-1, 1).to(device, dtype=torch.float32) + cond_latents.append(self.vae_encoder.encode(pix)) + return presentation, cond_latents, aligned_frames + + def _image_to_uint8_nhwc(self, img: torch.Tensor) -> torch.Tensor: + """Normalize a [1,3,H,W] image in either [-1,1] or [0,255] to [1,H,W,3] uint8.""" + x = img.float() + if x.numel() == 0: + return x + mx = x.max() + if mx > 2.0: + x = x / 255.0 # already 0..255 + else: + x = (x.clamp(-1, 1) + 1) * 0.5 # -1..1 -> 0..1 + return (x.movedim(1, -1).clamp(0, 1) * 255).to(torch.uint8) + @torch.inference_mode() def generate( self, @@ -146,6 +194,8 @@ class H3HotRuntime: steps: int, seed: int, attention: str | None = None, + first_frame: torch.Tensor | None = None, + last_frame: torch.Tensor | None = None, mux_audio: bool = True, ffmpeg_loglevel: str = "error", save_latent: str | Path | None = None, @@ -175,7 +225,24 @@ class H3HotRuntime: "latents_initialized", lambda: random_av_latents(width, height, frames, seed, device=self.config.device), ) - text = timed("text_conditioned", lambda: self.refiner(self.conditioner(prompt))) + use_fl2va = first_frame is not None or last_frame is not None + if use_fl2va: + # Build the fl2va Qwen presentation, encode the keyframes to cond latents, + # and refine the text span (vision-aware). + presentation, cond_latents, frame_count = timed( + "fl2va_conditioned", + lambda: self._build_fl2va(prompt, first_frame, last_frame, width, height, aligned_frames), + ) + text = timed("text_conditioned", lambda: self.refiner(presentation.text_states)) + pack_kwargs = { + "text_token_tags": presentation.text_token_tags, + "cond_latents": cond_latents, + "frame_count": frame_count, + "seed": seed, + } + else: + text = timed("text_conditioned", lambda: self.refiner(self.conditioner(prompt))) + pack_kwargs = {} sampled = timed( "sampled", lambda: sample_video_res_multistep( @@ -185,6 +252,7 @@ class H3HotRuntime: video, audio, steps=steps, + seed=seed, return_audio=mux_audio, cache_mode=cache_mode, cache_threshold=cache_threshold, @@ -192,6 +260,7 @@ class H3HotRuntime: cache_end_percent=cache_end_percent, cache_subsample_factor=cache_subsample_factor, cache_stats=cache_stats, + **pack_kwargs, ), ) if mux_audio: diff --git a/src/h3_blackwell_runtime/sampler.py b/src/h3_blackwell_runtime/sampler.py index 90a9e3e..ee029e9 100644 --- a/src/h3_blackwell_runtime/sampler.py +++ b/src/h3_blackwell_runtime/sampler.py @@ -77,6 +77,10 @@ def sample_video_res_multistep( model_timesteps: list[torch.Tensor] | tuple[torch.Tensor, ...] | None = None, return_audio: bool = False, progress: bool = False, + seed: int = 0, + text_token_tags: torch.Tensor | None = None, + cond_latents: list[torch.Tensor] | None = None, + frame_count: int | None = None, cache_mode: str | None = None, cache_threshold: float = 0.0, cache_start_percent: float = 0.0, @@ -136,7 +140,17 @@ def sample_video_res_multistep( carry = sigma_audio / sigma native_audio = audio_carried.to(torch.bfloat16) * carry step_timesteps = None if model_timesteps is None else model_timesteps[previous_index] - hidden, times, segments, positions, video_segment, audio_segment = packer(text, video, native_audio, _model_sigma(sigma), step_timesteps) + hidden, times, segments, positions, video_segment, audio_segment = packer( + text, + video, + native_audio, + _model_sigma(sigma), + step_timesteps, + text_token_tags=text_token_tags, + cond_latents=cond_latents, + frame_count=frame_count, + seed=seed, + ) raw_video, raw_audio = model(hidden, times, positions, segments, video_segment, audio_segment) raw_video = raw_video.to(torch.bfloat16).float() raw_audio = raw_audio.to(torch.bfloat16) diff --git a/src/h3_blackwell_runtime/vae_encoder.py b/src/h3_blackwell_runtime/vae_encoder.py new file mode 100644 index 0000000..a8d24bd --- /dev/null +++ b/src/h3_blackwell_runtime/vae_encoder.py @@ -0,0 +1,294 @@ +"""Direct, encoder-only MiniMax H3 video VAE implementation. + +Mirrors the encoder half of ``upstream_vae.py`` so keyframe/reference images can +be encoded without ComfyUI. The encoder runs in FP32 and the latent moments are +upcast to FP32 before mean/std normalization (the reference contract). + +Causal-conv semantics: spatial padding is reflect; temporal padding is causal +(front-only zeros) with a stride grid that starts at the first input frame +(Comfy autopad "same" / ``causal``). For a single input frame the temporal +taps of the kernel are truncated (Comfy ``autopad="causal_zero"``) so the frame +is not convolved against zero frames. +""" + +from __future__ import annotations + +import math +import os +from pathlib import Path + +import torch +from safetensors import safe_open +from torch import nn +from torch.nn import functional as F + + +IMAGENET_MEAN = (0.485, 0.456, 0.406) +IMAGENET_STD = (0.229, 0.224, 0.225) +LATENTS_MEAN = (0.858090341091156, -0.9606591463088989, 1.0661640167236328, -0.5090325474739075, -0.2727581858634949, -1.3675414323806763, -0.2553254961967468, -0.26907554268836975, -0.5376840829849243, -0.0464097298681736, 0.6657370328903198, 0.19690127670764923, -0.5460608005523682, -0.4035342037677765, -0.23683024942874908, 0.25928452610969543, -0.30133944749832153, 0.211341992020607, -1.1206848621368408, 0.3581933379173279, -0.04225143790245056, 0.2604829967021942, 0.22864092886447906, 0.7056031823158264) +LATENTS_STD = (1.2223774194717407, 1.2767263650894165, 1.68317747116088865, 1.7549455165863037, 1.5636216402053833, 2.194143533706665, 0.96531379222869875, 1.05698859691619875, 0.841948926448822, 0.7729952931404114, 1.8955937623977661, 0.946841835975647, 0.7996809482574463, 0.44988900423049925, 0.7197399735450745, 0.69362932443618775, 2.961095094680786, 2.7694199085235595, 3.0496184825897215, 2.1088054180145265, 3.276226282119751, 3.1627357006073, 2.28168129920959475, 2.6127843856811525) + + +def _causal_front_padding(t_in: int, kernel: int, stride: int) -> int: + """Front-only zero padding matching the reference causal 3D conv. + + ``padding = kernel - 1 - (t_out - 1) * stride`` with ``t_out = ceil(t_in / + stride)`` (the reference computes the output length from the unpadded + input). The result is non-negative: with a stride-s grid and a (2k+1) + kernel, ``ceil(t/stride) >= 1 + (t-1) // stride`` for all t, so + ``(t_in - 1) % stride * stride >= (kernel - 1) % (2 * stride)``. + """ + t_out = math.ceil(t_in / stride) + return max(0, kernel - 1 - (t_out - 1) * stride) + + +class _CausalConv3d(nn.Module): + """3D conv: reflect spatial padding, causal (front-zero) temporal padding.""" + + def __init__(self, in_channels: int, out_channels: int, kernel_size: int, stride: int | tuple[int, int, int] = 1, spatial_padding: int = 0): + super().__init__() + self.kernel_size = kernel_size + self.stride = stride if isinstance(stride, tuple) else (stride, stride, stride) + self.spatial_padding = spatial_padding + self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride=self.stride, padding=(0, spatial_padding, spatial_padding)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + t = x.shape[2] + if t == 1: + # A single input frame never convolves against zero frames: the + # temporal taps are truncated (Comfy autopad="causal_zero"). + if self.spatial_padding > 0: + half = (self.kernel_size - 1) // 2 + kernel = self.conv.weight[:, :, half:half + 2 * self.spatial_padding + 1] + return F.conv3d(x, kernel, self.conv.bias, (1, self.stride[1], self.stride[2]), (0, self.spatial_padding, self.spatial_padding)) + return self.conv(x) + front = _causal_front_padding(t, self.kernel_size, self.stride[0]) + if front > 0: + x = F.pad(x, (0, 0, 0, 0, front, 0)) + return self.conv(x) + + +class TemporalIsolatedGroupNorm(nn.GroupNorm): + """GroupNorm with statistics computed per frame (time merged into batch).""" + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if x.dim() != 5: + return super().forward(x) + b, c, t, h, w = x.shape + x = x.permute(0, 2, 1, 3, 4).contiguous().view(b * t, c, 1, h, w) + x = super().forward(x) + return x.view(b, t, c, h, w).permute(0, 2, 1, 3, 4).contiguous() + + +def group_norm_3d(num_channels: int) -> TemporalIsolatedGroupNorm: + return TemporalIsolatedGroupNorm(num_groups=32, num_channels=num_channels, eps=1e-6, affine=True) + + +class Downsample3D(nn.Module): + def __init__(self, in_channels: int, out_channels: int, time_stride: int = 1, space_stride: int = 2): + super().__init__() + self.space_stride = space_stride + self.conv = _CausalConv3d(in_channels, out_channels, kernel_size=3, stride=(time_stride, space_stride, space_stride), spatial_padding=0) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.space_stride == 2: + x = F.pad(x, (0, 1, 0, 1, 0, 0), mode="reflect") + return self.conv(x) + + +class ResnetBlock3D(nn.Module): + def __init__(self, in_channels: int, out_channels: int | None = None): + super().__init__() + self.in_channels = in_channels + self.out_channels = in_channels if out_channels is None else out_channels + self.norm1 = group_norm_3d(in_channels) + self.norm2 = group_norm_3d(self.out_channels) + self.conv1 = _CausalConv3d(in_channels, self.out_channels, kernel_size=3, spatial_padding=1) + self.conv2 = _CausalConv3d(self.out_channels, self.out_channels, kernel_size=3, spatial_padding=1) + if self.in_channels != self.out_channels: + self.nin_shortcut = _CausalConv3d(self.in_channels, self.out_channels, kernel_size=1, spatial_padding=0) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + h = self.conv1(F.silu(self.norm1(x), inplace=True)) + h = self.conv2(F.silu(self.norm2(h), inplace=True)) + if self.in_channels != self.out_channels: + x = self.nin_shortcut(x) + return h.add_(x) + + +class EncoderFCN3D(nn.Module): + def __init__(self, ch: int, ch_mult: tuple[int, ...], space_down: tuple[int, ...], time_down: tuple[int, ...], num_res_blocks: int, in_channels: int, z_channels: int, double_z: bool = True): + super().__init__() + self.num_levels = len(ch_mult) + self.num_res_blocks = [num_res_blocks] * self.num_levels + block_mid = [ch * ch_mult[i] for i in range(self.num_levels)] + block_in = [block_mid[0]] + block_mid[:-1] + self.conv_in = _CausalConv3d(in_channels, block_in[0], kernel_size=3, spatial_padding=1) + self.down = nn.ModuleList() + for i_level in range(self.num_levels): + down = nn.Module() + down.block = nn.ModuleList( + ResnetBlock3D(block_in[i_level] if i == 0 else block_mid[i_level], block_mid[i_level]) + for i in range(self.num_res_blocks[i_level]) + ) + if space_down[i_level] * time_down[i_level] > 1: + down.downsample = Downsample3D(block_mid[i_level], block_mid[i_level], time_stride=time_down[i_level], space_stride=space_down[i_level]) + self.down.append(down) + self.norm_out = group_norm_3d(block_mid[-1]) + self.conv_out = _CausalConv3d(block_mid[-1], 2 * z_channels if double_z else z_channels, kernel_size=3, spatial_padding=1) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + h = self.conv_in(x) + for i_level in range(self.num_levels): + for i_block in range(self.num_res_blocks[i_level]): + h = self.down[i_level].block[i_block](h) + if hasattr(self.down[i_level], "downsample"): + h = self.down[i_level].downsample(h) + h = F.silu(self.norm_out(h), inplace=True) + return self.conv_out(h) + + +class MiniMaxH3VideoVAEEncoder(nn.Module): + """Encoder-only H3 VAE. ``encode`` matches the public contract of upstream_vae.py.""" + + def __init__(self, *, device: torch.device | str | None = None, tiling: bool = True): + super().__init__() + self.vae_ratio, self.vae_ratio_t = 16, 4 + self.clip_length, self.token_drop = 17, 3 + self.tiling, self.tile_size, self.tile_overlap_min = tiling, 256, 64 + self.encoder = EncoderFCN3D(ch=128, ch_mult=(1, 2, 2, 4, 4, 8), space_down=(2, 2, 2, 2, 1, 1), time_down=(1, 2, 2, 1, 1, 1), num_res_blocks=2, in_channels=3, z_channels=24, double_z=True) + self.quant_conv = nn.Conv3d(48, 48, 1, device=device) + self.register_buffer("latents_mean", torch.tensor(LATENTS_MEAN), persistent=False) + self.register_buffer("latents_std", torch.tensor(LATENTS_STD), persistent=False) + self.register_buffer("pixel_mean", torch.tensor(IMAGENET_MEAN).view(1, 3, 1, 1, 1), persistent=False) + self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD).view(1, 3, 1, 1, 1), persistent=False) + + @classmethod + def from_safetensors(cls, path: str | Path, *, device: str | torch.device = "cuda", tiling: bool = True) -> "MiniMaxH3VideoVAEEncoder": + model = cls(device="meta", tiling=tiling) + expected = model.state_dict() + if os.getenv("H3_FAST_SAFETENSORS", "").lower() in {"1", "true", "yes", "on"}: + from fastsafetensors import fastsafe_open + + fast_device = "cuda:0" if str(device) == "cuda" else str(device) + available = set() + with fastsafe_open(filenames=[str(path)], nogds=True, device=fast_device) as checkpoint: + available = set(checkpoint.keys()) + missing = sorted(set(expected) - available) + weights = {name: checkpoint.get_tensor(name).clone().detach() for name in expected} + if missing: + raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(missing)}") + weights = {name: t.to(device=device, dtype=torch.float32) for name, t in weights.items()} + elif os.getenv("H3_DISABLE_MMAP", "").lower() in {"1", "true", "yes", "on"}: + from safetensors.torch import load + + with open(path, "rb") as file: + available_weights = load(file.read()) + available = set(available_weights) + missing = sorted(set(expected) - available) + if missing: + raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(missing)}") + weights = {name: available_weights[name].to(device=device, dtype=torch.float32) for name in expected} + else: + with safe_open(str(path), framework="pt", device=str(device)) as checkpoint: + available = set(checkpoint.keys()) + shape_errors = [(name, tuple(expected[name].shape), tuple(checkpoint.get_slice(name).get_shape())) for name in expected if name in available and tuple(expected[name].shape) != tuple(checkpoint.get_slice(name).get_shape())] + if shape_errors: + raise ValueError(f"incompatible H3 VAE checkpoint; shape mismatch: {shape_errors}") + weights = {name: checkpoint.get_tensor(name).to(dtype=torch.float32) for name in expected} + missing = sorted(set(expected) - available) + if missing: + raise ValueError(f"incompatible H3 VAE checkpoint; missing: {', '.join(missing)}") + model.load_state_dict(weights, strict=True, assign=True) + return model + + def _encode_moments(self, x: torch.Tensor) -> torch.Tensor: + return self.quant_conv(self.encoder(x)) + + def _adaptive_encode(self, x: torch.Tensor) -> torch.Tensor: + if self.tiling: + return self.tiled_encode(x) + return self._encode_moments(x) + + def split_tiles(self, length: int) -> tuple[list[int], list[int], list[int]]: + if self.tile_size >= length: + return [0], [length], [] + count = math.ceil(length / self.tile_size) + while self.tile_size * count - self.tile_overlap_min * (count - 1) < length: + count += 1 + overlaps = [self.tile_overlap_min] * (count - 1) + for index in range((self.tile_size * count - sum(overlaps) - length) // self.vae_ratio): + overlaps[index % len(overlaps)] += self.vae_ratio + starts = [0] + for overlap in overlaps: + starts.append(starts[-1] + self.tile_size - overlap) + return starts, [self.tile_size] * count, overlaps + + @staticmethod + def blend(a: torch.Tensor, b: torch.Tensor, extent: int, dim: int) -> torch.Tensor: + extent = min(a.shape[dim], b.shape[dim], extent) + positions = torch.arange(extent, device=b.device, dtype=b.dtype) + weight_a = 1 - positions / extent + weight_b = positions / extent + shape = [1] * a.ndim + shape[dim] = extent + weight_a = weight_a.view(shape) + weight_b = weight_b.view(shape) + slice_a = [slice(None)] * a.ndim + slice_a[dim] = slice(-extent, None) + slice_b = [slice(None)] * a.ndim + slice_b[dim] = slice(0, extent) + blended = a[tuple(slice_a)] * weight_a + b[tuple(slice_b)] * weight_b + if extent < b.shape[dim]: + slice_b_rest = [slice(None)] * b.ndim + slice_b_rest[dim] = slice(extent, None) + return torch.cat((blended, b[tuple(slice_b_rest)]), dim=dim) + return blended + + def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: + height, width = x.shape[-2], x.shape[-1] + y_idx, y_len, y_overlap = self.split_tiles(height) + x_idx, x_len, x_overlap = self.split_tiles(width) + rows = [[self._encode_moments(x[..., i_pos:i_pos + i_len, j_pos:j_pos + j_len]) for j_pos, j_len in zip(x_idx, x_len)] for i_pos, i_len in zip(y_idx, y_len)] + latent_y_overlap = [o // self.vae_ratio for o in y_overlap] + latent_x_overlap = [o // self.vae_ratio for o in x_overlap] + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + if i > 0: + tile = self.blend(rows[i - 1][j], tile, latent_y_overlap[i - 1], dim=-2) + if j > 0: + tile = self.blend(row[j - 1], tile, latent_x_overlap[j - 1], dim=-1) + if i < len(rows) - 1: + tile = tile[..., :-latent_y_overlap[i], :] + if j < len(row) - 1: + tile = tile[..., :, :-latent_x_overlap[j]] + result_row.append(tile) + result_rows.append(torch.cat(result_row, dim=-1)) + return torch.cat(result_rows, dim=-2) + + def encode_temporal(self, x: torch.Tensor) -> torch.Tensor: + if x.shape[2] % self.clip_length != 0: + pad_size = (-x.shape[2]) % self.clip_length + x = torch.cat([x, x[:, :, -1:].repeat(1, 1, pad_size, 1, 1)], dim=2) + num_chunks = x.shape[2] // self.clip_length + z_list = [self._adaptive_encode(x[:, :, i * self.clip_length:(i + 1) * self.clip_length, :, :]) for i in range(num_chunks)] + z = torch.cat(z_list, dim=2) + if self.token_drop > 0: + z = z[:, :, :-self.token_drop] + return z + + def encode(self, x: torch.Tensor) -> torch.Tensor: + """``[B, 3, T, H, W]`` pixels in ``[-1, 1]`` -> normalized latents ``[B, 24, T_lat, H//16, W//16]``.""" + if x.ndim == 4: + x = x.unsqueeze(2) + x = x.add(1.0).mul_(0.5).sub_(self.pixel_mean.to(x)).div_(self.pixel_std.to(x)) + if x.shape[2] == 1: + moments = self._adaptive_encode(x)[:, :, -1:, :, :] + else: + moments = self.encode_temporal(x) + mean = torch.chunk(moments.float(), 2, dim=1)[0] + latents_mean = self.latents_mean.view(1, -1, 1, 1, 1).to(mean) + latents_std = self.latents_std.view(1, -1, 1, 1, 1).to(mean) + return (mean - latents_mean) / latents_std diff --git a/tools/direct_t2v_preview.py b/tools/direct_t2v_preview.py index 49e313e..62793ac 100644 --- a/tools/direct_t2v_preview.py +++ b/tools/direct_t2v_preview.py @@ -7,6 +7,7 @@ import subprocess import time from datetime import datetime, timezone import warnings +import numpy as np warnings.filterwarnings("ignore", message="Found GPU0 NVIDIA GB10 which is of cuda capability 12.1.*", category=UserWarning) @@ -21,8 +22,10 @@ from h3_blackwell_runtime.packing import H3PromptPacker from h3_blackwell_runtime.qwen3vl_text import Qwen3VLPromptConditioner from h3_blackwell_runtime.sampler import sample_video_res_multistep from h3_blackwell_runtime.t2v import random_av_latents +from h3_blackwell_runtime.qwen3vl_vision import build_fl2va_presentation, resize_keyframe from h3_blackwell_runtime.token_refiner import H3TokenRefiner from h3_blackwell_runtime.vae_decoder import MiniMaxH3VideoVAE, dtype_from_name +from h3_blackwell_runtime.vae_encoder import MiniMaxH3VideoVAEEncoder parser = argparse.ArgumentParser() @@ -33,6 +36,8 @@ parser.add_argument("--height", type=int, default=192) parser.add_argument("--frames", type=int, default=22) parser.add_argument("--steps", type=int, default=12) parser.add_argument("--seed", type=int, default=440204) +parser.add_argument("--first-frame", type=Path, help="First keyframe image (fl2va), PNG/JPG path.") +parser.add_argument("--last-frame", type=Path, help="Last keyframe image (fl2va), PNG/JPG path.") parser.add_argument("--attention", choices=AVAILABLE_BACKENDS, default=DEFAULT_ATTENTION_BACKEND) parser.add_argument("--model-timesteps-capture", type=Path, help="Directory containing captured input_XX.pt H3 timesteps for strict parity checks.") parser.add_argument("--progress", action="store_true", help="Print per-step sampler progress.") @@ -97,8 +102,6 @@ video, audio, frames = random_av_latents(args.width, args.height, args.frames, a model = H3PackedDenoiser.from_checkpoint(checkpoint, attention_backend=args.attention).eval() configure_mlp_chunking(model, args.mlp_chunks, args.mlp_chunk_threshold) report_memory("h3_loaded") -text = H3TokenRefiner(checkpoint, attention_backend=args.attention)(conditioner(args.prompt)) -report_memory("text_conditioned") model_timesteps = None if args.model_timesteps_capture is not None: model_timesteps = [ @@ -107,23 +110,88 @@ if args.model_timesteps_capture is not None: ] want_audio = args.save_audio_latent is not None or args.audio_output is not None or args.mux_audio cache_stats = {} -sampled = sample_video_res_multistep( - model, - H3PromptPacker(checkpoint), - text, - video, - audio, - steps=args.steps, - model_timesteps=model_timesteps, - return_audio=want_audio, - progress=args.progress, - cache_mode=args.cache_mode, - cache_threshold=args.cache_threshold, - cache_start_percent=args.cache_start_percent, - cache_end_percent=args.cache_end_percent, - cache_subsample_factor=args.cache_subsample_factor, - cache_stats=cache_stats, -) +refiner = H3TokenRefiner(checkpoint, attention_backend=args.attention) +packer = H3PromptPacker(checkpoint) + +if args.first_frame is not None or args.last_frame is not None: + from PIL import Image + + def load_image(path: Path) -> torch.Tensor: + img = Image.open(path).convert("RGB") + return torch.from_numpy(np.array(img)).permute(2, 0, 1).unsqueeze(0).float() / 255.0 + + first = load_image(args.first_frame) if args.first_frame is not None else None + last = load_image(args.last_frame) if args.last_frame is not None else None + vae_encoder = MiniMaxH3VideoVAEEncoder.from_safetensors("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda").eval() + report_memory("vae_encoder_loaded") + from h3_blackwell_runtime.qwen3vl_vision import Qwen3VL32BVision + + vision_tower = Qwen3VL32BVision( + "/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors", device="cuda", dtype=torch.bfloat16 + ) + report_memory("vision_tower_loaded") + presentation = build_fl2va_presentation( + args.prompt, + first, + last, + width=args.width, + height=args.height, + frame_count=frames, + tokenizer=conditioner.tokenizer, + vision=vision_tower, + text_encoder=conditioner.encoder, + device="cuda", + ) + cond_latents = [] + for kf in presentation.keyframes: + resized = resize_keyframe(kf["image"].cuda(), args.width, args.height, crop="disabled" if kf["resolved_frame_index"] == 0 else "center") + cond_latents.append(vae_encoder.encode(resized.movedim(-1, 1).cuda().float())) + report_memory("fl2va_conditioned") + text = refiner(presentation.text_states) + report_memory("text_conditioned") + seed = args.seed + sampled = sample_video_res_multistep( + model, + packer, + text, + video, + audio, + steps=args.steps, + model_timesteps=model_timesteps, + return_audio=want_audio, + progress=args.progress, + seed=seed, + text_token_tags=presentation.text_token_tags, + cond_latents=cond_latents, + frame_count=frames, + cache_mode=args.cache_mode, + cache_threshold=args.cache_threshold, + cache_start_percent=args.cache_start_percent, + cache_end_percent=args.cache_end_percent, + cache_subsample_factor=args.cache_subsample_factor, + cache_stats=cache_stats, + ) +else: + text = refiner(conditioner(args.prompt)) + report_memory("text_conditioned") + sampled = sample_video_res_multistep( + model, + packer, + text, + video, + audio, + steps=args.steps, + model_timesteps=model_timesteps, + return_audio=want_audio, + progress=args.progress, + seed=args.seed, + cache_mode=args.cache_mode, + cache_threshold=args.cache_threshold, + cache_start_percent=args.cache_start_percent, + cache_end_percent=args.cache_end_percent, + cache_subsample_factor=args.cache_subsample_factor, + cache_stats=cache_stats, + ) if cache_stats: report({"cache": cache_stats}) if want_audio: diff --git a/tools/serve_hot_runtime.py b/tools/serve_hot_runtime.py index 2391131..a296f8b 100644 --- a/tools/serve_hot_runtime.py +++ b/tools/serve_hot_runtime.py @@ -10,10 +10,44 @@ from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from urllib.parse import urlparse +import torch +from PIL import Image + from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS, DEFAULT_ATTENTION_BACKEND from h3_blackwell_runtime.runtime import H3HotRuntime, RuntimeConfig +def _load_image(value) -> torch.Tensor | None: + """Accept a keyframe as on-disk path or base64 JPEG/PNG -> ``[1,3,H,W]`` float ``[0,1]``.""" + if value in (None, ""): + return None + if isinstance(value, (list, tuple)): + value = value[0] + if isinstance(value, dict): + value = value.get("url") or value.get("path") or value.get("b64") + if isinstance(value, str) and "\n" not in value and len(value) < 2048 and not value.startswith("data:"): + path = Path(value) + if path.exists(): + image = Image.open(path).convert("RGB") + import numpy as np + + tensor = torch.from_numpy(np.array(image))[None].permute(0, 3, 1, 2).float() / 255.0 + return tensor + data = value + if isinstance(value, str) and value.startswith("data:"): + data = value.split(",", 1)[1] + if isinstance(data, str) and len(data) >= 1024: + import base64 + import io + + raw = base64.b64decode(data) + image = Image.open(io.BytesIO(raw)).convert("RGB") + import numpy as np + + return torch.from_numpy(np.array(image))[None].permute(0, 3, 1, 2).float() / 255.0 + raise ValueError("first_frame/last_frame must be a path or a base64/data-URL image") + + parser = argparse.ArgumentParser() parser.add_argument("--host", default="0.0.0.0") parser.add_argument("--port", type=int, default=8000) @@ -91,6 +125,8 @@ class Handler(BaseHTTPRequestHandler): return mux_audio = bool(payload.get("mux_audio", True)) ffmpeg_loglevel = payload.get("ffmpeg_loglevel", "error") + first_frame = _load_image(payload.get("first_frame")) + last_frame = _load_image(payload.get("last_frame")) save_latent = payload.get("save_latent") cache_mode = payload.get("cache_mode") cache_threshold = float(payload.get("cache_threshold", 0.0)) @@ -108,6 +144,8 @@ class Handler(BaseHTTPRequestHandler): steps=steps, seed=seed, attention=attention, + first_frame=first_frame, + last_frame=last_frame, mux_audio=mux_audio, ffmpeg_loglevel=ffmpeg_loglevel, save_latent=save_latent,