Added additional information on addine new models and some additional gotchas

This commit is contained in:
Jaret Burkett
2026-06-18 15:04:50 -06:00
parent 515b0ea5cd
commit e886745051
3 changed files with 64 additions and 0 deletions

View File

@@ -90,6 +90,11 @@ train:
(`get_loss_target`), i.e. the velocity pointing from data to noise.
- `self.model` / `self.transformer` / `self.unet` are aliases for the same
thing on BaseModel.
- **`use_old_lokr_format = False`** — set this class attribute on every NEW
model. `BaseModel` defaults it to `True` purely for backwards-compatibility
with LoKr checkpoints trained before the format change; all new architectures
should use the new LoKr format. (Plain LoRA training is unaffected — this only
matters for `network.type: "lokr"`.)
## AdvancedPromptEmbeds
@@ -145,6 +150,51 @@ a `Conv3d` you don't actually run — e.g. a vision tower's patch embed on a VL
text encoder — drop it (`text_encoder.model.visual = None`) to skip loading it;
if you must run one, consider running that component in fp16/fp32.
## Attention backends (don't force flash-attn)
Reference repos very often hard-code an attention kernel — `flash_attn`,
xformers, sage — and import it at module top level. **Do not carry that
requirement over.** ai-toolkit has to import and load your model on machines
where that package isn't installed (CPU boxes, headless CI, plain installs), so
a top-level `from flash_attn import ...` turns "load the model" into an
`ImportError`.
The rule:
- **Default to torch's built-in `F.scaled_dot_product_attention`** (the
"native" backend). It needs no extra dependency, runs on CPU and CUDA, and
already dispatches to a fused/flash kernel on supported hardware. `src/model.py`
does exactly this.
- **Make any other kernel OPTIONAL**, selected at runtime — never required at
import. The clean pattern:
1. Guard the import so a missing package is a flag, not a crash:
```python
try:
from flash_attn import flash_attn_varlen_func
_FLASH_ATTN_AVAILABLE = True
except ImportError:
flash_attn_varlen_func = None
_FLASH_ATTN_AVAILABLE = False
```
2. Give each attention module an `attention_backend` flag (default
`"native"`) and **branch inside its forward** — `"flash"` runs the flash
kernel, anything else runs SDPA.
3. Expose a `set_attention_backend("native"|"flash")` on the parent model
that validates the name, raises a clear error if `"flash"` is requested
while `_FLASH_ATTN_AVAILABLE` is `False`, and propagates the flag to every
attention module.
4. Wire it to a config knob so it stays opt-in, e.g.
`model_kwargs.attention_backend: "flash"` read in `load_model`.
Branch on a per-module **flag**, don't swap the processor/module instance:
attention modules that own trained q/k/v weights (joint/dual-stream blocks)
would lose those weights if you replaced them with a different instance.
Worked implementations to copy: `../ideogram4/src/transformer.py`
(`set_attention_backend`, native+flash in one `Attention.forward`) and
`../boogu_image/src/attention_processor.py` (guarded import, per-processor
`attention_backend` flag, flash varlen branch alongside SDPA).
## Adapting this template
### Editing / instruct model (image in, image out)

View File

@@ -56,6 +56,11 @@ class ExampleModel(BaseModel):
# - it is the default cache key for text-embedding / latent caches
arch = "example"
# ALL NEW MODELS should set this to False. ``BaseModel`` defaults it to True
# only for backwards-compatibility with already-released LoKr checkpoints; the
# newer LoKr weight format is the correct one for any new architecture.
use_old_lokr_format = False
def __init__(
self,
device, # "cuda:0" etc.

View File

@@ -111,6 +111,15 @@ class ExampleTransformerBlock(nn.Module):
) # each (B, 1, hidden), broadcasts over the sequence
# --- attention ---
# ALWAYS default to torch's built-in scaled_dot_product_attention so the
# model runs with no extra dependency. If the reference repo you are
# porting hard-codes flash-attn (or xformers, sage, ...), do NOT carry
# that requirement over -- make it OPTIONAL. The clean pattern is a
# per-module ``attention_backend`` flag toggled in bulk from the parent
# model (e.g. ``set_attention_backend("flash")``), branching to the
# flash kernel only when explicitly selected AND the package is present.
# See ../../ideogram4/src/transformer.py and ../../boogu_image/src for
# working "native" (SDPA) + optional "flash" implementations.
h = self.norm1(x) * (1 + scale_a) + shift_a
q, k, v = self.qkv(h).chunk(3, dim=-1)
q = q.view(b, s, self.num_heads, self.head_dim).transpose(1, 2)