diff --git a/README.md b/README.md index 5e3141e..5b7c432 100644 --- a/README.md +++ b/README.md @@ -59,3 +59,8 @@ docker compose -f compose.spark.yml up -d h3-hot-runtime ``` Use `GET /ready` to confirm resident model readiness. Use `POST /generate` with JSON fields like `prompt`, `output`, `width`, `height`, `frames`, `steps`, `seed`, and optional `attention`. Supported request-level attention values are reported by `/ready`; switching attention does not reload model weights. + +Exact memory/lifetime options: + +- `attention: "kj_head_sliced"` slices attention heads and runs the slice backend from `H3_HEAD_SLICE_BACKEND` (`sage2` by default) with `H3_HEAD_SLICE_SIZE` heads per slice (`8` by default). +- `--mlp-chunks N` on `tools/serve_hot_runtime.py` or `tools/direct_t2v_preview.py` chunks H3 SwiGLU rows exactly to reduce peak activation memory. Default is `1` (disabled). diff --git a/src/h3_blackwell_runtime/attention.py b/src/h3_blackwell_runtime/attention.py index bf570db..f892665 100644 --- a/src/h3_blackwell_runtime/attention.py +++ b/src/h3_blackwell_runtime/attention.py @@ -1,5 +1,6 @@ """Direct H3 self-attention using packed NVFP4 linears and SageAttention3.""" +import os import torch import torch.nn.functional as functional from torch import nn @@ -8,8 +9,8 @@ from .checkpoint import H3Checkpoint from .nvfp4 import Nvfp4Linear -AVAILABLE_BACKENDS = ("sage2", "sdpa", "sage3", "sage3_mean", "kj_sage_cuda", "kj_sage_triton", "kj_sage_fp8", "kj_sage_fp8pp") -PLANNED_BACKENDS = ("flash4", "easycache", "h3_cache", "sol_attn", "kj_chunked_ffn", "kj_head_sliced") +AVAILABLE_BACKENDS = ("sage2", "sdpa", "sage3", "sage3_mean", "kj_sage_cuda", "kj_sage_triton", "kj_sage_fp8", "kj_sage_fp8pp", "kj_head_sliced") +PLANNED_BACKENDS = ("flash4", "easycache", "h3_cache", "sol_attn", "kj_chunked_ffn") def attention_backend_status() -> dict[str, str]: @@ -19,8 +20,7 @@ def attention_backend_status() -> dict[str, str]: status.update({"easycache": "planned: approximate denoiser cache"}) status.update({"h3_cache": "planned: approximate H3-specific cache"}) status.update({"sol_attn": "experimental: prior H3-tested sparse Triton attention; source/package pending"}) - status.update({"kj_chunked_ffn": "planned: exact memory-lifetime adapter"}) - status.update({"kj_head_sliced": "planned: exact memory-lifetime adapter"}) + status.update({"kj_chunked_ffn": "available: exact H3 MLP row chunking via H3_MLP_CHUNKS or runtime args"}) return status @@ -30,6 +30,16 @@ def rms_norm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor: def run_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *, backend: str, is_causal: bool) -> torch.Tensor: """Run one `[batch, heads, sequence, dim]` attention operation.""" + if backend == "kj_head_sliced": + head_slice_size = int(os.getenv("H3_HEAD_SLICE_SIZE", "8")) + base_backend = os.getenv("H3_HEAD_SLICE_BACKEND", "sage2") + if head_slice_size <= 0: + raise ValueError("H3_HEAD_SLICE_SIZE must be positive") + outputs = [ + run_attention(q[:, start:start + head_slice_size], k[:, start:start + head_slice_size], v[:, start:start + head_slice_size], backend=base_backend, is_causal=is_causal) + for start in range(0, q.shape[1], head_slice_size) + ] + return torch.cat(outputs, dim=1) if backend == "sage2": from sageattention import sageattn diff --git a/src/h3_blackwell_runtime/block.py b/src/h3_blackwell_runtime/block.py index 2ff7622..7e27fa8 100644 --- a/src/h3_blackwell_runtime/block.py +++ b/src/h3_blackwell_runtime/block.py @@ -39,6 +39,8 @@ class H3SwiGLU(nn.Module): super().__init__() self.fc1 = fc1 self.fc2 = fc2 + self.chunks = 1 + self.chunk_threshold = 4096 @classmethod def from_checkpoint(cls, checkpoint: H3Checkpoint, prefix: str, *, output_dtype=torch.bfloat16): @@ -48,10 +50,25 @@ class H3SwiGLU(nn.Module): ) def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.chunks > 1 and x.shape[0] >= self.chunk_threshold: + return torch.cat([self._forward_chunk(chunk) for chunk in x.chunk(self.chunks, dim=0)], dim=0) + return self._forward_chunk(x) + + def _forward_chunk(self, x: torch.Tensor) -> torch.Tensor: gate, up = self.fc1(x).chunk(2, dim=-1) return self.fc2(torch.nn.functional.silu(gate).mul_(up)) +def configure_mlp_chunking(model: nn.Module, chunks: int, threshold: int = 4096) -> None: + """Configure exact row-chunked H3 SwiGLU execution to reduce peak activation memory.""" + if chunks < 1: + raise ValueError("MLP chunks must be >= 1") + for module in model.modules(): + if isinstance(module, H3SwiGLU): + module.chunks = chunks + module.chunk_threshold = threshold + + class H3DiTBlock(nn.Module): """One H3 transformer block with externally supplied AdaLN tensors.""" diff --git a/src/h3_blackwell_runtime/runtime.py b/src/h3_blackwell_runtime/runtime.py index c2756a0..ac3a638 100644 --- a/src/h3_blackwell_runtime/runtime.py +++ b/src/h3_blackwell_runtime/runtime.py @@ -11,6 +11,7 @@ import torch from .audio_vae_decoder import MiniMaxH3AudioVAE from .attention import AVAILABLE_BACKENDS +from .block import configure_mlp_chunking from .checkpoint import H3Checkpoint from .denoiser import H3PackedDenoiser from .packing import H3PromptPacker @@ -31,6 +32,8 @@ class RuntimeConfig: attention: str = "sage2" vae_dtype: str = "float16" vae_tile_size: int = 256 + mlp_chunks: int = 1 + mlp_chunk_threshold: int = 4096 device: str = "cuda" @@ -59,7 +62,7 @@ class H3HotRuntime: ) self.model = self._timed_load( "h3_loaded", - lambda: H3PackedDenoiser.from_checkpoint(self.checkpoint, attention_backend=config.attention).eval(), + self._load_h3, ) self.refiner = self._timed_load( "token_refiner_loaded", @@ -92,6 +95,11 @@ class H3HotRuntime: vae.tile_size = self.config.vae_tile_size return vae + def _load_h3(self) -> H3PackedDenoiser: + model = H3PackedDenoiser.from_checkpoint(self.checkpoint, attention_backend=self.config.attention).eval() + configure_mlp_chunking(model, self.config.mlp_chunks, self.config.mlp_chunk_threshold) + return model + def status(self) -> dict: return { "ready": True, @@ -99,6 +107,8 @@ class H3HotRuntime: "current_attention": self.attention, "vae_dtype": self.config.vae_dtype, "vae_tile_size": self.config.vae_tile_size, + "mlp_chunks": self.config.mlp_chunks, + "mlp_chunk_threshold": self.config.mlp_chunk_threshold, "loaded_at": self.loaded_at, "load_stages": self.load_stages, } diff --git a/tools/direct_t2v_preview.py b/tools/direct_t2v_preview.py index 1c031f6..f9d2ae4 100644 --- a/tools/direct_t2v_preview.py +++ b/tools/direct_t2v_preview.py @@ -15,6 +15,7 @@ import torch from h3_blackwell_runtime.checkpoint import H3Checkpoint from h3_blackwell_runtime.audio_vae_decoder import MiniMaxH3AudioVAE from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS +from h3_blackwell_runtime.block import configure_mlp_chunking from h3_blackwell_runtime.denoiser import H3PackedDenoiser from h3_blackwell_runtime.packing import H3PromptPacker from h3_blackwell_runtime.qwen3vl_text import Qwen3VLPromptConditioner @@ -45,6 +46,8 @@ parser.add_argument("--mux-audio", action="store_true") parser.add_argument("--skip-decode", action="store_true") parser.add_argument("--vae-dtype", choices=("float32", "float16", "bfloat16"), default=os.getenv("H3_VAE_DTYPE", "float16")) parser.add_argument("--vae-tile-size", type=int, default=int(os.getenv("H3_VAE_TILE_SIZE", "256"))) +parser.add_argument("--mlp-chunks", type=int, default=int(os.getenv("H3_MLP_CHUNKS", "1"))) +parser.add_argument("--mlp-chunk-threshold", type=int, default=int(os.getenv("H3_MLP_CHUNK_THRESHOLD", "4096"))) args = parser.parse_args() started = time.perf_counter() last_report = started @@ -87,6 +90,7 @@ conditioner = Qwen3VLPromptConditioner( report_memory("qwen_loaded") video, audio, frames = random_av_latents(args.width, args.height, args.frames, args.seed) 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") diff --git a/tools/serve_hot_runtime.py b/tools/serve_hot_runtime.py index a620809..fd12624 100644 --- a/tools/serve_hot_runtime.py +++ b/tools/serve_hot_runtime.py @@ -20,11 +20,13 @@ parser.add_argument("--port", type=int, default=8000) parser.add_argument("--attention", choices=AVAILABLE_BACKENDS, default="sage2", help="Initial attention backend. Requests can switch with the JSON attention field.") parser.add_argument("--vae-dtype", choices=("float32", "float16", "bfloat16"), default="float16") parser.add_argument("--vae-tile-size", type=int, default=256) +parser.add_argument("--mlp-chunks", type=int, default=1) +parser.add_argument("--mlp-chunk-threshold", type=int, default=4096) parser.add_argument("--warmup", action="store_true", help="Run a tiny generation before accepting traffic.") parser.add_argument("--warmup-output", type=Path, default=Path("/output/h3-blackwell-runtime/hot-runtime-warmup.mp4")) args = parser.parse_args() -runtime = H3HotRuntime(RuntimeConfig(attention=args.attention, vae_dtype=args.vae_dtype, vae_tile_size=args.vae_tile_size)) +runtime = H3HotRuntime(RuntimeConfig(attention=args.attention, vae_dtype=args.vae_dtype, vae_tile_size=args.vae_tile_size, mlp_chunks=args.mlp_chunks, mlp_chunk_threshold=args.mlp_chunk_threshold)) runtime_lock = threading.Lock() warmup_result = None if args.warmup: