Add Sol-Attn backend
This commit is contained in:
parent
75ee9ba4ca
commit
3d0c093168
4 changed files with 53 additions and 5 deletions
|
|
@ -1,6 +1,8 @@
|
|||
# GB10/Grace Blackwell development image. It does not start inference by default.
|
||||
FROM ghcr.io/aeon-7/comfyui-aeon-spark:slim
|
||||
|
||||
ARG SOL_ATTN_COMMIT=930a4d6e432ff8b8ed5e30ff2f72519b92d69bdf
|
||||
|
||||
WORKDIR /opt/h3-blackwell-runtime
|
||||
COPY . .
|
||||
|
||||
|
|
@ -13,10 +15,18 @@ RUN python -m pip install --no-cache-dir --no-deps comfy-kitchen==0.2.28
|
|||
|
||||
RUN python -m pip install --no-cache-dir "fastsafetensors>=0.1.10"
|
||||
|
||||
RUN python -m pip uninstall -y pynvml \
|
||||
&& python -m pip install --no-cache-dir nvidia-ml-py
|
||||
|
||||
RUN git clone https://github.com/Saganaki22/ComfyUI-sol-attn.git /opt/ComfyUI-sol-attn \
|
||||
&& cd /opt/ComfyUI-sol-attn \
|
||||
&& git checkout ${SOL_ATTN_COMMIT}
|
||||
|
||||
RUN python -m pip install --no-cache-dir --no-deps -e . \
|
||||
&& python -c "import comfy_kitchen, torch; from sageattn3 import sageattn3_blackwell; assert hasattr(torch.ops.comfy_kitchen, 'rms_rope_split_half_'); print(torch.__version__, torch.version.cuda)"
|
||||
|
||||
ENV H3_MODEL_PATH=/models/minimax_h3_ref2va_pruned_nvfp4.safetensors
|
||||
ENV PYTHONPATH=/opt/ComfyUI-sol-attn
|
||||
ENV TORCH_COMPILE_DISABLE=0 TORCHDYNAMO_DISABLE=0
|
||||
ENTRYPOINT []
|
||||
CMD ["bash"]
|
||||
|
|
|
|||
|
|
@ -63,6 +63,7 @@ Use `GET /ready` to confirm resident model readiness. Use `POST /generate` with
|
|||
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).
|
||||
- `attention: "sol_attn"` routes eligible H3 attention calls through the pinned ComfyUI Sol-Attn Triton kernel vendored into the Spark image. Configure with `H3_SOL_TAU` (`1.3`), `H3_SOL_MIN_TOKENS` (`4096`), `H3_SOL_THRESH_TYPE` (`diag`), `H3_SOL_INT8_QK`, `H3_SOL_INT8_PV`, `H3_SOL_FALLBACK` (`sage2`), and `H3_SOL_STRICT`.
|
||||
- `--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).
|
||||
|
||||
Approximate cache options are opt-in and must be quality-gated per prompt:
|
||||
|
|
|
|||
|
|
@ -9,17 +9,17 @@ 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", "kj_head_sliced")
|
||||
PLANNED_BACKENDS = ("flash4", "easycache", "h3_cache", "sol_attn", "kj_chunked_ffn")
|
||||
AVAILABLE_BACKENDS = ("sage2", "sdpa", "sage3", "sage3_mean", "kj_sage_cuda", "kj_sage_triton", "kj_sage_fp8", "kj_sage_fp8pp", "kj_head_sliced", "sol_attn")
|
||||
PLANNED_BACKENDS = ("flash4", "easycache", "h3_cache", "kj_chunked_ffn")
|
||||
|
||||
|
||||
def attention_backend_status() -> dict[str, str]:
|
||||
"""Report direct-runtime attention choices without importing ComfyUI nodes."""
|
||||
status = {name: "available" for name in AVAILABLE_BACKENDS}
|
||||
status.update({"sol_attn": "experimental: sparse Triton attention for eligible non-causal H3 attention calls; falls back below H3_SOL_MIN_TOKENS unless H3_SOL_STRICT=1"})
|
||||
status.update({"flash4": "planned: exact Blackwell kernel adapter"})
|
||||
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": "available: exact H3 MLP row chunking via H3_MLP_CHUNKS or runtime args"})
|
||||
return status
|
||||
|
||||
|
|
@ -30,6 +30,35 @@ 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 == "sol_attn":
|
||||
try:
|
||||
from sol_kernel import sol_attn
|
||||
|
||||
tau = float(os.getenv("H3_SOL_TAU", "1.3"))
|
||||
min_tokens = int(os.getenv("H3_SOL_MIN_TOKENS", "4096"))
|
||||
thresh_type = os.getenv("H3_SOL_THRESH_TYPE", "diag")
|
||||
int8_qk = os.getenv("H3_SOL_INT8_QK", "").lower() in {"1", "true", "yes", "on"}
|
||||
int8_pv = os.getenv("H3_SOL_INT8_PV", "").lower() in {"1", "true", "yes", "on"}
|
||||
if is_causal:
|
||||
raise ValueError("Sol-Attn backend only supports non-causal H3 attention")
|
||||
if q.shape[-1] != 128:
|
||||
raise ValueError(f"Sol-Attn requires head dim 128, got {q.shape[-1]}")
|
||||
if q.shape[2] < min_tokens:
|
||||
raise ValueError(f"{q.shape[2]} tokens < H3_SOL_MIN_TOKENS={min_tokens}")
|
||||
out = sol_attn(
|
||||
q.transpose(1, 2).contiguous(),
|
||||
k.transpose(1, 2).contiguous(),
|
||||
v.transpose(1, 2).contiguous(),
|
||||
tau=tau,
|
||||
thresh_type=thresh_type,
|
||||
int8_qk=int8_qk,
|
||||
int8_pv=int8_pv,
|
||||
)
|
||||
return out.transpose(1, 2)
|
||||
except Exception:
|
||||
if os.getenv("H3_SOL_STRICT", "").lower() in {"1", "true", "yes", "on"}:
|
||||
raise
|
||||
return run_attention(q, k, v, backend=os.getenv("H3_SOL_FALLBACK", "sage2"), is_causal=is_causal)
|
||||
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")
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -46,6 +47,13 @@ def _ffmpeg_command(loglevel: str, *parts: str) -> list[str]:
|
|||
return ["ffmpeg", "-hide_banner", "-loglevel", loglevel, *parts]
|
||||
|
||||
|
||||
def _refiner_attention_backend(attention: str) -> str:
|
||||
if attention != "sol_attn":
|
||||
return attention
|
||||
fallback = os.getenv("H3_SOL_FALLBACK", "sage2")
|
||||
return fallback if fallback in AVAILABLE_BACKENDS and fallback != "sol_attn" else "sage2"
|
||||
|
||||
|
||||
class H3HotRuntime:
|
||||
"""Keep all prompt-only H3 models resident for repeated requests."""
|
||||
|
||||
|
|
@ -66,7 +74,7 @@ class H3HotRuntime:
|
|||
)
|
||||
self.refiner = self._timed_load(
|
||||
"token_refiner_loaded",
|
||||
lambda: H3TokenRefiner(self.checkpoint, attention_backend=config.attention).eval(),
|
||||
lambda: H3TokenRefiner(self.checkpoint, attention_backend=_refiner_attention_backend(config.attention)).eval(),
|
||||
)
|
||||
self.packer = H3PromptPacker(self.checkpoint)
|
||||
self.video_vae = self._timed_load(
|
||||
|
|
@ -123,7 +131,7 @@ class H3HotRuntime:
|
|||
if hasattr(module, "backend"):
|
||||
module.backend = attention
|
||||
for block in self.refiner.blocks:
|
||||
block.attention_backend = attention
|
||||
block.attention_backend = _refiner_attention_backend(attention)
|
||||
self.attention = attention
|
||||
|
||||
@torch.inference_mode()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue