This commit is contained in:
Jaret Burkett
2026-08-27 16:16:35 -06:00
parent 9113420b61
commit 45886f01b2
2 changed files with 31 additions and 2 deletions

View File

@@ -382,8 +382,12 @@ Decisions:
raw `quantize()` call sites and silently dropped by `quantize_model()`; the
ARA path inside `quantize_model` hardcodes `uint8`. Decide the unified
behavior when consolidating.
- [ ] `toolkit/models/loaders/umt5.py` dead `comfy_files` param (wan21 comfy-TE
no-op) — fix when wan migrates.
- [x] wan comfy-TE resolved for real: UMT5TextEncoder carries comfy
candidates (fp8_e4m3fn_scaled via the legacy importer, fp16), files
already in transformers key layout (spiece blob dropped, tied
embed_tokens materialized). Verified: wan21 samples with the local
comfy fp8 TE. The loaders/umt5.py `comfy_files` param stays as a
no-op shim for old callers.
- [ ] Registry hardening: error (don't fall back to SD1) on unknown arch; lazy
per-arch imports; single source of truth shared with the UI's
`options.tsx` model list.

View File

@@ -33,11 +33,36 @@ class UMT5TextEncoder(UMT5EncoderModel, OstrisTransformersMixin):
aitk_subfolder = "text_encoder"
aitk_tokenizer_subfolder = "tokenizer"
aitk_config_repo = "ai-toolkit/umt5_xxl_encoder"
aitk_comfy_repo = "Comfy-Org/Wan_2.1_ComfyUI_repackaged"
# comfy umt5 files already use the transformers key layout; the
# fp8_e4m3fn_scaled variant is the legacy scaled-fp8 format (handled by
# the importer's float8 backend)
aitk_comfy_weight_names = {
"ai-toolkit/umt5_xxl_encoder": [
"split_files/text_encoders/umt5_xxl_fp8_e4m3fn_scaled.safetensors",
"split_files/text_encoders/umt5_xxl_fp16.safetensors",
],
}
@classmethod
def get_transformer_block_names(cls):
return ["encoder.block"]
@classmethod
def convert_state_dict_on_load(cls, state_dict):
# drop the embedded sentencepiece blob and materialize the tied
# embed_tokens reference
state_dict = dict(state_dict)
state_dict.pop("spiece_model", None)
if (
"encoder.embed_tokens.weight" not in state_dict
and "shared.weight" in state_dict
):
state_dict["encoder.embed_tokens.weight"] = state_dict["shared.weight"]
return state_dict
@classmethod
def load_tokenizer(cls, name_or_path=None, subfolder=None, **kwargs):
# T5's tokenizer needs the _spm_precompiled_charsmap patch