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.
|
# GB10/Grace Blackwell development image. It does not start inference by default.
|
||||||
FROM ghcr.io/aeon-7/comfyui-aeon-spark:slim
|
FROM ghcr.io/aeon-7/comfyui-aeon-spark:slim
|
||||||
|
|
||||||
|
ARG SOL_ATTN_COMMIT=930a4d6e432ff8b8ed5e30ff2f72519b92d69bdf
|
||||||
|
|
||||||
WORKDIR /opt/h3-blackwell-runtime
|
WORKDIR /opt/h3-blackwell-runtime
|
||||||
COPY . .
|
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 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 . \
|
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)"
|
&& 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 H3_MODEL_PATH=/models/minimax_h3_ref2va_pruned_nvfp4.safetensors
|
||||||
|
ENV PYTHONPATH=/opt/ComfyUI-sol-attn
|
||||||
ENV TORCH_COMPILE_DISABLE=0 TORCHDYNAMO_DISABLE=0
|
ENV TORCH_COMPILE_DISABLE=0 TORCHDYNAMO_DISABLE=0
|
||||||
ENTRYPOINT []
|
ENTRYPOINT []
|
||||||
CMD ["bash"]
|
CMD ["bash"]
|
||||||
|
|
|
||||||
|
|
@ -63,6 +63,7 @@ Use `GET /ready` to confirm resident model readiness. Use `POST /generate` with
|
||||||
Exact memory/lifetime options:
|
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: "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).
|
- `--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:
|
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
|
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")
|
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", "sol_attn", "kj_chunked_ffn")
|
PLANNED_BACKENDS = ("flash4", "easycache", "h3_cache", "kj_chunked_ffn")
|
||||||
|
|
||||||
|
|
||||||
def attention_backend_status() -> dict[str, str]:
|
def attention_backend_status() -> dict[str, str]:
|
||||||
"""Report direct-runtime attention choices without importing ComfyUI nodes."""
|
"""Report direct-runtime attention choices without importing ComfyUI nodes."""
|
||||||
status = {name: "available" for name in AVAILABLE_BACKENDS}
|
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({"flash4": "planned: exact Blackwell kernel adapter"})
|
||||||
status.update({"easycache": "planned: approximate denoiser cache"})
|
status.update({"easycache": "planned: approximate denoiser cache"})
|
||||||
status.update({"h3_cache": "planned: approximate H3-specific 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"})
|
status.update({"kj_chunked_ffn": "available: exact H3 MLP row chunking via H3_MLP_CHUNKS or runtime args"})
|
||||||
return status
|
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:
|
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."""
|
"""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":
|
if backend == "kj_head_sliced":
|
||||||
head_slice_size = int(os.getenv("H3_HEAD_SLICE_SIZE", "8"))
|
head_slice_size = int(os.getenv("H3_HEAD_SLICE_SIZE", "8"))
|
||||||
base_backend = os.getenv("H3_HEAD_SLICE_BACKEND", "sage2")
|
base_backend = os.getenv("H3_HEAD_SLICE_BACKEND", "sage2")
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
import subprocess
|
import subprocess
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
@ -46,6 +47,13 @@ def _ffmpeg_command(loglevel: str, *parts: str) -> list[str]:
|
||||||
return ["ffmpeg", "-hide_banner", "-loglevel", loglevel, *parts]
|
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:
|
class H3HotRuntime:
|
||||||
"""Keep all prompt-only H3 models resident for repeated requests."""
|
"""Keep all prompt-only H3 models resident for repeated requests."""
|
||||||
|
|
||||||
|
|
@ -66,7 +74,7 @@ class H3HotRuntime:
|
||||||
)
|
)
|
||||||
self.refiner = self._timed_load(
|
self.refiner = self._timed_load(
|
||||||
"token_refiner_loaded",
|
"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.packer = H3PromptPacker(self.checkpoint)
|
||||||
self.video_vae = self._timed_load(
|
self.video_vae = self._timed_load(
|
||||||
|
|
@ -123,7 +131,7 @@ class H3HotRuntime:
|
||||||
if hasattr(module, "backend"):
|
if hasattr(module, "backend"):
|
||||||
module.backend = attention
|
module.backend = attention
|
||||||
for block in self.refiner.blocks:
|
for block in self.refiner.blocks:
|
||||||
block.attention_backend = attention
|
block.attention_backend = _refiner_attention_backend(attention)
|
||||||
self.attention = attention
|
self.attention = attention
|
||||||
|
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue