Models v2 - phase 0

This commit is contained in:
Jaret Burkett
2026-08-27 10:51:43 -06:00
parent 5497a001cb
commit e8d9cf6d35
29 changed files with 583 additions and 401 deletions

View File

@@ -98,6 +98,9 @@ UNET_IN_CHANNELS = 4 # Stable Diffusion の in_channels は 4 で固定。XLも
class BaseModel:
# override these in child classes
arch = None
# rename LoRA keys transformer. <-> diffusion_model. (the ComfyUI-standard
# prefix) on save/load
lora_keys_use_comfy_prefix = False
def __init__(
self,
@@ -1627,10 +1630,20 @@ class BaseModel:
def convert_lora_weights_before_save(self, state_dict):
# can be overridden in child classes to convert weights before saving
if self.lora_keys_use_comfy_prefix:
return {
k.replace("transformer.", "diffusion_model."): v
for k, v in state_dict.items()
}
return state_dict
def convert_lora_weights_before_load(self, state_dict):
# can be overridden in child classes to convert weights before loading
if self.lora_keys_use_comfy_prefix:
return {
k.replace("diffusion_model.", "transformer."): v
for k, v in state_dict.items()
}
return state_dict
def condition_noisy_latents(self, latents: torch.Tensor, batch:'DataLoaderBatchDTO'):

View File

@@ -0,0 +1,227 @@
# v2 Model Module Restructure — Planning
## Goal
Every component the toolkit loads (DiTs/transformers, unets, text encoders, vision
encoders, VAEs, audio VAEs) becomes a class extending one base module in
`toolkit/models/v2`. `BaseModel` (`toolkit/models/base_model.py`) stays as the
multimodal holder that each arch in `extensions_built_in/diffusion_models` extends —
that layer is good. The layer below it is what gets unified: one loading entry point,
one quantization path, one save path, shared component definitions instead of
per-model-folder copies.
End state this enables:
- **Live server with model hot-swap**: a resident process where, when a generation or
training run requests a different model, the unused components are dropped and the
new ones loaded. Shared component classes (same TE/VAE reused across archs) make
component-level reuse possible instead of full teardown/reload.
- **Model loading test suite**: a test that loads each registered arch one at a time
and runs an inference pass. Every model type gets added to this suite as it is
migrated (see Testing below).
- **Comfy-aligned weights**: weights live in the ComfyUI folder layout under
`MODELS_PATH` (shareable with a ComfyUI install), download there when missing, and
saves are comfy-format. Eventually defaults move to comfy / our own prequantized
releases for everything.
## Current state (survey 2026-08-27)
Three generations of loading conventions coexist:
1. Legacy monolith `toolkit/stable_diffusion_model.py` (`is_flux` / `is_v3` branches);
still the silent fallback in `toolkit/util/get_model.py` when an arch string
doesn't match.
2. `BaseModel` subclass per arch (35 registered classes), each with a hand-written
`load_model()` / `save_model()`.
3. `toolkit/models/v2/_mixin.py` (`OstrisModelMixin`) — the intended fix, currently
used by one model (`v2/z_image.py` → z_image extension).
### Duplication highlights
- BFL KL autoencoder: full copies in `flux2/src/autoencoder.py` and
`ideogram4/src/vae.py` (header says "Flux2 KL autoencoder").
- Qwen3-VL text encoder loaded independently in qwen_image, nucleus_image, krea2,
ideogram4, mageflow, minimax_h3 — from several different repo sources; krea2
hand-patches the vision tower locally.
- `Qwen3ForCausalLM` TE load: 3 verbatim line-for-line copies
(`z_image/z_image.py:257`, `z_image/z_image_l2p_model.py:471`,
`zeta_chroma/zeta_chroma_model.py:146`).
- Flux1 VAE + T5 + CLIP trio loaded 4x from 2 different repos (chroma ×2,
flux_kontext, legacy SD path).
- Comfy-file resolver copy-pasted: `minimax_h3/minimax_h3.py:248`
(`_resolve_comfy_file`) → `ltx2/ltx2.py:1254` ("mirrors MinimaxH3Model").
- `AutoencoderKLQwenImage` latents mean/std handling triplicated (qwen_image,
nucleus_image, krea2).
- `transformer.` ↔ `diffusion_model.` LoRA key rename copy-pasted ~15x in
`convert_lora_weights_before_save/load` overrides.
- Fake CLIP/TE/config stubs redefined in ~5 places
(canonical: `toolkit/models/FakeVAE.py`, `toolkit/unloader.py`).
### Inconsistency highlights
- **Quantization, 5 paths**: `quantize_model()` (block-streaming, ARA-aware — ~21
users), raw `quantize()` (~10 users, no block streaming/excludes, but the only path
honoring `quantize_kwargs`), hidream's hand-rolled block loop, the v2 mixin's own
`quantize_`, and bare TE quantization everywhere.
- **Known bugs**: ~9 sites quantize the TE with `qtype` instead of `qtype_te`
(chroma ×2, flux2, flux_kontext, cogview4, wan21, legacy SD, ...);
`toolkit/models/loaders/umt5.py` accepts a `comfy_files` param it never uses, so
wan21's comfy-TE path is a silent no-op.
- **Saving, 4 incompatible styles**: diffusers `save_pretrained` folders, flat
safetensors, safetensors-inside-diffusers-folder hybrids, and z_image's
loaded-format-dependent branch. Dequant-on-save done 3 ways; the
`isinstance(v, QTensor)` variant (chroma, flux2, boogu_image, ideogram4) misses
torchao and Ostris weights entirely. Every `save_pretrained` override ignores its
`save_dtype` argument.
- Registry: linear scan in `toolkit/util/get_model.py`, silent SD1 fallback on a
typo'd arch, eager import of every model file at startup. Second unsynchronized
registry in `ui/src/app/jobs/new/options.tsx`.
## Decisions (locked in)
1. **Save format = ComfyUI format.** Single-file safetensors in comfy key layout.
Must support saving quantized — primarily convrot8 and nvfp4 (comfy_quant marker
format, see `toolkit/util/comfy_quant_import.py`) — and plain bf16, all in comfy
format. Loading stays backwards compatible: diffusers dirs, transformers repos,
and single files all still digest through `load_model`; only saving standardizes
on comfy.
2. **v2 folder layout mirrors the comfy save path structure**:
```
toolkit/models/v2/
_mixin.py # base module (OstrisModelMixin, evolving)
resolver.py # comfy-layout weight resolution (lift from minimax_h3)
diffusion_models/ # one file per DiT/unet family
text_encoders/ # qwen3_vl.py, qwen3.py, t5.py, clip.py, gemma.py, ...
vae/ # flux_kl.py, qwen_image.py, wan.py, audio VAEs, ...
vision_encoders/
```
3. **Method names win from `BaseModel`**: `get_transformer_block_names` and
`get_quantization_exclude_modules`. The mixin's `get_quantization_block_names`
gets renamed to match; resolve the classmethod-vs-instance-method mismatch while
doing so.
4. **Loading policy**:
- Per-model special handling is allowed via the hook methods.
- If `name_or_path` is a diffusers/transformers source, load it with
diffusers/transformers for now. **Step 1 is migrating every model to the v2
module format and loader without breaking anything** — same weights, same
sources, same results.
- Each model declares a `comfy_weight_names` dict keyed per standard
`name_or_path`. If the user points at a local folder or a non-standard repo,
load it as-is. If `name_or_path` is the standard repo and we have matching
comfy weight names, load those instead when any of them exist (locally under
`MODELS_PATH` in comfy layout, or downloadable to there).
- Eventually the default flips to comfy weights / our own prequantized releases
for everything.
## Base module: what `OstrisModelMixin` still needs
The mixin already handles: diffusers dir / hub repo / local single file /
`org/repo/file.safetensors`, key-conversion hooks on load and save, overridable
backend hooks for transformers-lib models, block-wise quantize.
To add:
- [x] **Comfy weight spec + resolver.** `aitk_comfy_repo` / `aitk_comfy_weight_names`
class attrs + `find_comfy_weights` (local-only until Phase 2); resolution chain
generalized from `minimax_h3._resolve_comfy_file` into `v2/resolver.py`:
explicit override → `MODELS_PATH` at the repo-relative comfy path → flat at
root → recursive walk of the category folder → hub download **to the
repo-relative path** (folder stays shareable with ComfyUI, no duplicate
downloads).
- [x] **Automatic prequantized import.** Single-file path sniffs `comfy_quant`
markers and routes through `import_comfy_quantized_layers` before
`load_state_dict`, including the OstrisLinear missing-key whitelist that
minimax_h3 and ltx2 each hand-rolled.
- [x] **One save path.** `save_model(path, dtype)`: dequantize via
`dequantize_if_quantized` (honors dtype), run `convert_state_dict_on_save`,
write single-file comfy-layout safetensors. (Quantized-storage saves —
convrot8 / nvfp4 with comfy_quant markers — land with Phase 2; diffusers-folder
save as an explicit flag still to add.)
- [x] **Tokenizer/processor declaration** for text encoders
(`aitk_tokenizer_repo`/`aitk_processor_repo` + `load_tokenizer`/`load_processor`).
- [x] Rename quantization hooks to the `BaseModel` spellings (decision 3):
`get_transformer_block_names` (classmethod on the module).
## Migration steps
Track progress here; check items off as they land.
### Phase 0 — foundation (done 2026-08-27)
- [x] Evolve `_mixin.py` per the list above (comfy spec local-only until Phase 2)
- [x] Create `v2/diffusion_models/`, `v2/text_encoders/`, `v2/vae/`,
`v2/vision_encoders/`; move `v2/z_image.py` → `v2/diffusion_models/z_image.py`
- [x] Lift the comfy resolver out of minimax_h3 into `v2/resolver.py`; point
minimax_h3 and ltx2 at it (delete their copies)
- [x] `BaseModel` default `convert_lora_weights_before_save/load` doing the
`transformer.` ↔ `diffusion_model.` rename, gated on the class attr
`lora_keys_use_comfy_prefix` (default False, so passthrough models keep
their behavior); the ~18 identical overrides replaced with the flag.
Custom conversions (anima, hidream_o1, ltx2, wan21) keep their overrides;
ltx2's now composes with the flag via super().
### Phase 1 — migrate all models to v2 modules, no behavior change
Every arch's components become v2 classes; if `name_or_path` is diffusers, it still
loads via diffusers. Nothing about sources or outputs changes yet. Suggested order
(worst duplication first), each including its loading test (see Testing):
- [ ] Shared TEs: `text_encoders/qwen3_vl.py`, `text_encoders/qwen3.py`,
`text_encoders/t5.py`, `text_encoders/clip.py`
- [ ] Shared VAEs: `vae/flux_kl.py` (kills flux2 + ideogram4 copies and the 4
scattered `AutoencoderKL` loads), `vae/qwen_image.py` (mean/std handling
built in)
- [ ] z_image / z_image_l2p (already half-migrated)
- [ ] qwen_image family (qwen_image, qwen_image_edit, qwen_image_edit_plus)
- [ ] nucleus_image, krea2, ideogram4, mageflow
- [ ] chroma, chroma_radiance, zeta_chroma
- [ ] flux2, flux_kontext
- [ ] minimax_h3 (+ ref2va), ltx2 family
- [ ] wan21 / wan22 family
- [ ] hidream family, omnigen2
- [ ] anima, boogu_image, ernie_image, f_light, prx_pixel_t2i
- [ ] audio_models (ace_step)
- [ ] Per-model fixes folded in as each migrates: `qtype_te` bug, dequant-on-save
(`dequantize_if_quantized` everywhere), raw-`quantize()` → `quantize_model()`
### Phase 2 — comfy weights become the preferred source
- [ ] Wire `comfy_weight_names` per model; standard-repo `name_or_path` + existing
comfy weights → load comfy
- [ ] Comfy-format save (bf16 + convrot8/nvfp4 quantized) as the default
full-weight save
- [ ] Publish/verify comfy repacks per model as they flip
### Phase 3 — live server
- [ ] Component-level identity (which TE/VAE instances are shared between archs) so
a model switch drops only what the next run doesn't need
- [ ] Resident process: request comes in → diff requested components vs loaded →
unload/load the difference
- [ ] Legacy `stable_diffusion_model.py` archs: grandfather or port last
## Testing
- [ ] `testing/` (or `tests/`) harness: for each migrated arch, load the model via
its v2 modules and run one small inference pass (single low-step sample; video
models at minimum frame count). One arch at a time, full unload between archs.
- [ ] Weights resolved through the normal resolver against `MODELS_PATH`
(GPU + local-weights test, not CI-portable at first; skip archs whose weights
are absent rather than failing).
- [ ] Round-trip test per model: load → save comfy format → reload from the save →
outputs match (bf16) / load cleanly (quantized saves).
- [ ] Each newly migrated model adds its test in the same PR as its migration.
## TODO / look at later
- [ ] Quantize-path consolidation quirks: `quantize_kwargs` is honored only by the
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.
- [ ] 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.
- [ ] Fake/stub components: consolidate on `toolkit/models/FakeVAE.py` /
`toolkit/unloader.py`, delete local copies.
- [ ] Vendored upstream code (hidream/src, omnigen2/src, ltx2 converter's private
comfy-quant parser): dedupe against toolkit utils where practical.

View File

@@ -9,7 +9,11 @@ digests any of:
- a HuggingFace repo id ("org/repo")
- a local single .safetensors file in the model's original key layout
- a remote single .safetensors file ("org/repo/path/file.safetensors")
plus optional automatic quantization of the loaded weights (`qtype=...`).
plus optional automatic quantization of the loaded weights (`qtype=...`). A
single-file checkpoint carrying ComfyUI `comfy_quant` markers loads with its
pre-quantized layers attached to the toolkit's quantization backends
(convrot8 / nvfp4) automatically. `save_model` writes a single-file
.safetensors back in the original (comfy) key layout.
Model specific behavior lives in small overridable hooks (key conversion, config
source, block names, backend loading), so subclassing for a new model usually means
@@ -47,6 +51,19 @@ class OstrisModelMixin:
# .safetensors file and no config_path is given
aitk_config_repo: Optional[str] = None
# ---- comfy weight sources (Phase 2 flips the loading default to these) ----
# hub repo the comfy-format weight files are published in
aitk_comfy_repo: Optional[str] = None
# standard name_or_path (hub repo id) -> repo-relative comfy weight file
# (ComfyUI folder layout: diffusion_models/, text_encoders/, vae/, ...)
aitk_comfy_weight_names: Dict[str, str] = {}
# ---- tokenizer/processor source, for text-encoder modules ----
aitk_tokenizer_repo: Optional[str] = None
aitk_tokenizer_subfolder: Optional[str] = None
aitk_processor_repo: Optional[str] = None
aitk_processor_subfolder: Optional[str] = None
# ---- state set by the loader / quantizer ----
aitk_is_quantized: bool = False
aitk_qtype: Optional[str] = None
@@ -68,7 +85,7 @@ class OstrisModelMixin:
return state_dict
@classmethod
def get_quantization_block_names(cls) -> Optional[List[str]]:
def get_transformer_block_names(cls) -> Optional[List[str]]:
"""Names (dotted paths allowed) of the repeated block lists to quantize one
block at a time so the whole model never has to sit on the gpu at once."""
return None
@@ -200,16 +217,132 @@ class OstrisModelMixin:
state_dict = load_file(file_path)
state_dict = cls.convert_state_dict_on_load(state_dict)
for key, value in state_dict.items():
state_dict[key] = value.to(dtype=dtype)
has_quant_markers = any(k.endswith(".comfy_quant") for k in state_dict)
model = cls.aitk_from_config(config)
model.load_state_dict(state_dict, assign=True)
model.to(dtype=dtype)
if has_quant_markers:
# pre-quantized comfy checkpoint: attach the quantized linears onto
# the toolkit's backends, load the rest at its stored precision
# (the checkpoint's bf16/fp16/fp32 mix is deliberate)
from toolkit.util.comfy_quant_import import import_comfy_quantized_layers
state_dict, num_quantized = import_comfy_quantized_layers(
model, state_dict, orig_dtype=dtype
)
cls._load_state_dict_with_quantized(model, state_dict)
model.aitk_is_quantized = True
else:
for key, value in state_dict.items():
state_dict[key] = value.to(dtype=dtype)
model.load_state_dict(state_dict, assign=True)
model.to(dtype=dtype)
del state_dict
flush()
return model
@staticmethod
def _load_state_dict_with_quantized(model, state_dict):
"""Load a state dict onto a (meta-built) model whose quantized linears
were already attached by import_comfy_quantized_layers: their weight
(and importer-assigned bias) legitimately report as missing keys."""
from toolkit.util.ostris_quant import OstrisLinear
result = model.load_state_dict(state_dict, assign=True, strict=False)
quantized_param_keys = set()
for name, m in model.named_modules():
if isinstance(m, OstrisLinear):
quantized_param_keys.add(f"{name}.weight")
if m.bias is not None:
quantized_param_keys.add(f"{name}.bias")
bad_missing = [k for k in result.missing_keys if k not in quantized_param_keys]
if bad_missing or result.unexpected_keys:
raise ValueError(
f"{type(model).__name__} load mismatch: missing {bad_missing[:8]}, "
f"unexpected {result.unexpected_keys[:8]}"
)
leftover_meta = [n for n, p in model.named_parameters() if p.is_meta]
if leftover_meta:
raise ValueError(
f"{type(model).__name__} load left meta parameters: "
f"{leftover_meta[:8]}"
)
# ------------------------------------------------------------------
# comfy weight sources
# ------------------------------------------------------------------
@classmethod
def find_comfy_weights(cls, name_or_path: str) -> Optional[str]:
"""Local comfy-format weight file registered for a standard
``name_or_path``, or None. Never downloads — Phase 2 flips the loading
default to comfy sources; until then callers opt in explicitly."""
from toolkit.models.v2.resolver import resolve_comfy_file
rel_path = cls.aitk_comfy_weight_names.get(name_or_path)
if rel_path is None:
return None
return resolve_comfy_file(
rel_path, repo_id=cls.aitk_comfy_repo, local_only=True
)
# ------------------------------------------------------------------
# tokenizer / processor (text-encoder modules)
# ------------------------------------------------------------------
@classmethod
def load_tokenizer(cls, **kwargs):
from transformers import AutoTokenizer
if cls.aitk_tokenizer_repo is None:
raise ValueError(f"{cls.__name__} does not declare aitk_tokenizer_repo")
return AutoTokenizer.from_pretrained(
cls.aitk_tokenizer_repo, subfolder=cls.aitk_tokenizer_subfolder, **kwargs
)
@classmethod
def load_processor(cls, **kwargs):
from transformers import AutoProcessor
if cls.aitk_processor_repo is None:
raise ValueError(f"{cls.__name__} does not declare aitk_processor_repo")
return AutoProcessor.from_pretrained(
cls.aitk_processor_repo, subfolder=cls.aitk_processor_subfolder, **kwargs
)
# ------------------------------------------------------------------
# saving
# ------------------------------------------------------------------
@torch.no_grad()
def save_model(
self,
output_path: str,
dtype: Optional[torch.dtype] = None,
metadata: Optional[Dict[str, str]] = None,
):
"""Save as a single-file .safetensors in the model's original (comfy)
key layout, via convert_state_dict_on_save. Quantized weights are
dequantized to full precision; dtype, when given, casts the floating
point tensors. (Quantized-storage saves — comfy_quant markers — come
with Phase 2.)"""
from safetensors.torch import save_file
from toolkit.util.quantize import dequantize_if_quantized
state_dict = {}
for key, value in self.state_dict().items():
value = dequantize_if_quantized(value)
if dtype is not None and value.is_floating_point():
value = value.to(dtype=dtype)
state_dict[key] = value.detach().to("cpu").contiguous()
state_dict = self.convert_state_dict_on_save(state_dict)
parent = os.path.dirname(output_path)
if parent:
os.makedirs(parent, exist_ok=True)
save_file(state_dict, output_path, metadata=metadata)
del state_dict
flush()
# ------------------------------------------------------------------
# quantization
# ------------------------------------------------------------------
@@ -222,7 +355,7 @@ class OstrisModelMixin:
exclude: Optional[List[str]] = None,
):
"""Quantize the model weights in place. When device is given, the repeated
blocks (get_quantization_block_names) are moved there one at a time for the
blocks (get_transformer_block_names) are moved there one at a time for the
quantization math and returned to their original device, so the whole model
never has to fit on the gpu in full precision."""
from optimum.quanto import freeze
@@ -238,7 +371,7 @@ class OstrisModelMixin:
)
blocks: List[torch.nn.Module] = []
for name in self.get_quantization_block_names() or []:
for name in self.get_transformer_block_names() or []:
# name may be a dotted path for models that nest their blocks
block_list = self
for part in name.split("."):

View File

@@ -3,7 +3,7 @@ from diffusers.models.transformers import (
ZImageTransformer2DModel as DiffusersZImageTransformer2DModel,
)
from ._mixin import OstrisModelMixin
from .._mixin import OstrisModelMixin
class ZImageTransformer2DModel(DiffusersZImageTransformer2DModel, OstrisModelMixin):
@@ -12,7 +12,7 @@ class ZImageTransformer2DModel(DiffusersZImageTransformer2DModel, OstrisModelMix
aitk_config_repo = "Tongyi-MAI/Z-Image-Turbo"
@classmethod
def get_quantization_block_names(cls):
def get_transformer_block_names(cls):
return ["layers"]
@classmethod

View File

@@ -0,0 +1,129 @@
"""ComfyUI-layout weight file resolution.
Weight files live under MODELS_PATH in ComfyUI's folder layout
(diffusion_models/, text_encoders/, vae/, ...) so the folder is shareable with
a ComfyUI install. Files are used in place when present and downloaded to
exactly their repo-relative location only when missing, so nothing is ever
duplicated on re-run.
Lifted from the minimax_h3 / ltx2.5 model implementations; those now call
into here.
"""
import os
from typing import Callable, Iterable, Optional
from toolkit.paths import MODELS_PATH
def find_file_recursive(root_dir: str, filename: str) -> Optional[str]:
"""First (breadth-stable, sorted) match of ``filename`` anywhere under
``root_dir``."""
if not os.path.isdir(root_dir):
return None
for dirpath, dirnames, filenames in os.walk(root_dir):
dirnames.sort()
if filename in filenames:
return os.path.join(dirpath, filename)
return None
def repo_id_from_name_or_path(
name_or_path: Optional[str], default: str
) -> str:
"""Treat a hub-style ``name_or_path`` ("org/repo") as a replacement comfy
repo; anything local (or an explicit .safetensors file) keeps the
default."""
if (
name_or_path
and not os.path.exists(name_or_path)
and not name_or_path.endswith(".safetensors")
and "/" in name_or_path
):
return name_or_path
return default
def resolve_comfy_file(
rel_path: str,
repo_id: str,
override_path: Optional[str] = None,
extra_roots: Optional[Iterable[str]] = None,
hf_token: Optional[str] = None,
status_fn: Optional[Callable[[str], None]] = None,
local_only: bool = False,
) -> Optional[str]:
"""Find a weight file at its local location, or download it there when
(and only when) it is missing.
Search order: ``override_path`` (must exist), the repo-relative path under
MODELS_PATH (and each of ``extra_roots``), the bare filename at each root,
any subfolder of the category folder (recursive — e.g.
diffusion_models/my_custom_sub/), then the hub — downloaded to the
repo-relative path under MODELS_PATH. With ``local_only`` the hub is never
touched and a miss returns None.
"""
if override_path is not None:
if not os.path.exists(override_path):
raise FileNotFoundError(
f"Override path for {rel_path} does not exist: {override_path}"
)
return override_path
filename = os.path.basename(rel_path)
category = os.path.dirname(rel_path)
roots = [MODELS_PATH] + [r for r in (extra_roots or []) if os.path.isdir(r)]
for root in roots:
for rel in (rel_path, filename):
candidate = os.path.join(root, rel)
if os.path.exists(candidate):
return candidate
for root in roots:
found = find_file_recursive(os.path.join(root, category), filename)
if found is not None:
return found
if local_only:
return None
import huggingface_hub
if status_fn is not None:
status_fn(f"Downloading {rel_path} from {repo_id} into {MODELS_PATH}")
return huggingface_hub.hf_hub_download(
repo_id=repo_id, filename=rel_path, token=hf_token, local_dir=MODELS_PATH
)
def resolve_named_file(
path: str,
component: str = "model",
hf_token: Optional[str] = None,
) -> str:
"""Resolve an explicit .safetensors reference: a local file, a file already
under MODELS_PATH, or an 'org/repo/path/file.safetensors' hub path
(downloaded into the models folder at its repo-relative path)."""
if os.path.exists(path):
return path
splits = path.split("/")
if len(splits) < 3:
raise ValueError(
f"Invalid {component} path: {path}. Must be a local file or "
"'org/repo/filename.safetensors' to download from the Hugging Face Hub."
)
rel_path = "/".join(splits[2:])
for candidate in (
os.path.join(MODELS_PATH, rel_path),
os.path.join(MODELS_PATH, splits[-1]),
):
if os.path.exists(candidate):
return candidate
import huggingface_hub
return huggingface_hub.hf_hub_download(
repo_id="/".join(splits[:2]),
filename=rel_path,
token=hf_token,
local_dir=MODELS_PATH,
)

View File