MiniMax H3 Chunk FeedForward
The node that cuts H3's MLP peak by 37% without changing a single pixel
- model
- MODEL
Here's the shape of the problem. MiniMax H3's feed-forward starts with a projection to 28,672 features per token - 56 KB per token of float intermediate, before the swiglu. At an 8K packed sequence that's a chunk of VRAM you don't get back, and H3's joint packed sequence (text, conditioning, reference, audio, video) is long by design. On the pruned INT8 ConvRot checkpoint most 16GB cards run, that intermediate is where MLP peak memory goes: the activation function is fused into the quantizer rather than living as a separate pass.
MiniMax H3 Chunk FeedForward fixes that by running the MLP in token chunks instead of all at once. It's the H3 version of a trick people know from Kijai's Wan Chunk FeedForward node, and like that one it's the least glamorous way to buy VRAM back.
How it works
The MLP is token-local - no token talks to another token inside it - which is what makes chunking safe. The node replaces each MLP's forward with one that splits the activation along the sequence dimension, runs the original linear_input_act implementation on each chunk, and writes back into a preallocated output. Instead of holding one [tokens, 28672] intermediate, you hold roughly [tokens/chunks, 28672].
Because it preserves the original path inside each chunk, the INT8 ConvRot quantization with its row-wise activation scales comes out unchanged: the author measured outputs with torch.testing.assert_close at rtol=atol=0 and throughput neutral within ±4%. Not a trade - a straight swap of peak memory for a few extra launches.
The measured numbers, on a real first-block MLP loaded from minimax_h3_fl2va_pruned_int8_convrot.safetensors, at 2 chunks:
| Tokens | Full MLP | Chunked ×2 | Saved | |---:|---:|---:|---:| | 8,192 | 644 MiB | 406 MiB | 238 MiB | | 32,768 | 2,576 MiB | 1,624 MiB | 952 MiB | | 65,536 | 5,152 MiB | 3,248 MiB | 1,904 MiB |
A flat 37% off peak, and it scales with length. At short sequences the saving is real but small - which is exactly what min_tokens is for.
One important caveat, straight from the README: this is most effective on the INT8 ConvRot checkpoint. On a plain bf16 checkpoint the eager swiglu path already holds extra intermediates, so chunking buys you much less.
The inputs you actually touch
model- from your H3 loader.enabled- bypass. Note the node also returns the model untouched ifchunksis left at 1.chunks(2, range 1–64) - more chunks, lower peak, more overhead. If you're scraping against an OOM at 65K tokens, try 4.min_tokens(8192) - below this packed sequence length the normal full-width MLP runs and the node does nothing at all. Lower it if you want relief on shorter renders; raise it if you only care about long ones.
Output is a single MODEL into your guider.
It doesn't need Triton, and it doesn't care about attention
This is the one node in the pack that works on any attention backend - Sage, SDPA, Sol, or nothing patched at all. The pack's __init__.py guards the Triton node module and the MiniMax module independently, so the feed-forward node keeps loading even if the Sol kernel fails to import. If you've been holding off on the whole pack because Triton installation looked like a headache, this node is usable on its own.
It patches every DiT block's MLP, plus the token refiner blocks - you'll see something like [MiniMax H3 FFN] patched 62 MLPs (chunks=2, min_tokens=8192) in the console. Inputs that require gradients get the original unchunked MLP, since this is an inference optimization.
Install
cd ComfyUI/custom_nodes
git clone https://github.com/r-vage/ComfyUI-sol-attn
Restart ComfyUI, or install via ComfyUI Manager (publisher rvage, display name "ComfyUI Sol-Attn (continued by r-vage)"). pyproject.toml lists torch as the only hard dependency - no pip install needed if ComfyUI already runs H3.
Where people get burned
"Expected a MiniMax H3 model; returning it unchanged." You didn't hand it H3. It warns and does nothing rather than failing.
It has no effect on short clips. With min_tokens at its default of 8192, anything under that packed length skips the node entirely - and the packed sequence is longer than your frame count suggests, since audio, conditioning and reference tokens are in there too.
Chasing more chunks than you need. chunks 8 doesn't beat chunks 4 by much, and every split is extra launches. Go up a notch only when you're actually at the memory ceiling.
Pushing every memory lever at once. This one is free, so stack it with the pack's attention patches and a step-skipping cache like EasyCache all you like - just remember the approximations are elsewhere. Chunking changes nothing about your output; a high tau and an aggressive cache threshold definitely do.
Inputs (4)
| Name | Type | Default | Description |
|---|---|---|---|
| model | MODEL | — | |
| enabled | BOOLEAN | true | — |
| chunks | INT | 21–64 | More chunks reduce peak MLP activation memory but add overhead. |
| min_tokens | INT | 8192256–131072 | Keep the normal full-width MLP below this packed sequence length. |
Outputs (1)
| Name | Type | Description |
|---|---|---|
| MODEL | MODEL | — |