Spaces:
Running on Zero
Running on Zero
File size: 10,772 Bytes
c1e2af3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 | # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""Load Ardy diffusion models from local checkpoints or Hugging Face."""
from pathlib import Path
from typing import Optional
import torch
from huggingface_hub import snapshot_download
from omegaconf import OmegaConf
from .loading import (
DEFAULT_MODEL,
DEFAULT_TEXT_ENCODER_URL,
get_env_var,
instantiate_from_dict,
)
from .registry import hf_repo_id, resolve_model_name
DEFAULT_TEXT_ENCODER = "llm2vec"
TEXT_ENCODER_PRESETS = {
"llm2vec": {
"target": "ardy.model.LLM2VecEncoder",
"kwargs": {
# The original McGill adapter points at Meta's gated repository.
# Load the same Llama 3.0 architecture from an ungated mirror, then
# apply the original MNTP and supervised adapters explicitly.
"base_model_name_or_path": get_env_var(
"LLM2VEC_BASE_MODEL",
"unsloth/llama-3-8b-Instruct",
),
"mntp_model_name_or_path": get_env_var(
"LLM2VEC_MNTP_MODEL",
"McGill-NLP/LLM2Vec-Meta-Llama-3-8B-Instruct-mntp",
),
"peft_model_name_or_path": get_env_var(
"LLM2VEC_SUPERVISED_MODEL",
"McGill-NLP/LLM2Vec-Meta-Llama-3-8B-Instruct-mntp-supervised",
),
"prompt_model_family": "llama3-instruct",
"dtype": "bfloat16",
"llm_dim": 4096,
"device": "auto",
},
}
}
def _download_from_hf(full_name: str) -> Path:
"""Download a released model from Hugging Face; returns the local snapshot dir.
With LOCAL_CACHE=true, tries the local HF cache first and falls back online.
"""
repo_id = hf_repo_id(full_name)
local_cache = get_env_var("LOCAL_CACHE", "False").lower() == "true"
if local_cache:
try:
return Path(snapshot_download(repo_id=repo_id, local_files_only=True))
except Exception:
pass # cache miss -> download online below
return Path(snapshot_download(repo_id=repo_id))
def _build_api_text_encoder_conf(text_encoder_url: str) -> dict:
return {
"_target_": "ardy.model.text_encoder_api.TextEncoderAPI",
"url": text_encoder_url,
}
def _build_local_text_encoder_conf(text_encoder_fp32: bool = False) -> dict:
text_encoder_name = get_env_var("TEXT_ENCODER", DEFAULT_TEXT_ENCODER)
if text_encoder_name not in TEXT_ENCODER_PRESETS:
available = ", ".join(sorted(TEXT_ENCODER_PRESETS))
raise ValueError(f"Unknown TEXT_ENCODER='{text_encoder_name}'. Available: {available}")
preset = TEXT_ENCODER_PRESETS[text_encoder_name]
# Copy before overriding so the shared preset dict is never mutated.
kwargs = dict(preset["kwargs"])
if text_encoder_fp32:
kwargs["dtype"] = "float32"
return {
"_target_": preset["target"],
**kwargs,
}
def _select_text_encoder_conf(
text_encoder_url: str,
text_encoder_fp32: bool = False,
mode: str = "auto",
) -> tuple:
"""Return ``(conf, probe)``: the selected encoder conf plus the already-instantiated encoder
when auto-mode probing built one (reused by the caller so the API client is not constructed
twice).
``mode`` is resolved and validated by load_text_encoder:
- "api": force TextEncoderAPI
- "local": force local LLM2VecEncoder
- "auto": try API first, fallback to local if unreachable
"""
if mode == "local":
return _build_local_text_encoder_conf(text_encoder_fp32), None
if mode == "api":
return _build_api_text_encoder_conf(text_encoder_url), None
api_conf = _build_api_text_encoder_conf(text_encoder_url)
try:
text_encoder = instantiate_from_dict(api_conf)
# Probe availability early so inference doesn't fail later.
text_encoder(["healthcheck"])
return api_conf, text_encoder
except Exception as error:
print(
"Text encoder service is unreachable, falling back to local LLM2Vec "
f"encoder. ({type(error).__name__}: {error})"
)
return _build_local_text_encoder_conf(text_encoder_fp32), None
def load_text_encoder(
mode: Optional[str] = None,
url: Optional[str] = None,
fp32: bool = False,
device: Optional[str] = None,
):
"""Select and instantiate a text encoder, ready for inference.
This is the single place that owns text-encoder selection + instantiation,
so it can be built once and reused across multiple models (e.g. core / g1 /
soma) by passing the result into ``load_model(..., text_encoder=...)``.
Args:
mode: Backend selection ("auto"/"api"/"local"). When None, falls back
to the TEXT_ENCODER_MODE env var (default "auto").
url: Remote service URL. When None, falls back to the TEXT_ENCODER_URL
env var.
fp32: Use float32 instead of the default bfloat16.
device: Target device. When None, uses cuda if available else cpu.
Returns:
The instantiated text encoder placed on ``device``.
"""
if mode is None:
mode = get_env_var("TEXT_ENCODER_MODE", "auto")
mode = str(mode).lower()
if mode not in ("auto", "api", "local"):
raise ValueError(
f"Unknown text-encoder mode {mode!r}. Choose 'auto', 'api' or 'local'; "
"to load a model without a text encoder, pass text_encoder=False to "
"load_model()."
)
resolved_url = url or get_env_var("TEXT_ENCODER_URL", DEFAULT_TEXT_ENCODER_URL)
print(
f"Setting up text encoder (mode={mode}); first run may take a while...",
flush=True,
)
conf, text_encoder = _select_text_encoder_conf(resolved_url, fp32, mode=mode)
if text_encoder is None:
# Placement is handled below via .to(); drop any device kwarg the preset
# may carry so it doesn't conflict (e.g. accelerate device_map="auto").
conf = dict(conf)
conf.pop("device", None)
text_encoder = instantiate_from_dict(conf)
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.float32 if fp32 else torch.bfloat16
return text_encoder.to(device=device, dtype=dtype)
def load_model(
modelname=None,
device=None,
eval_mode: bool = True,
default_family: Optional[str] = None,
text_encoder=None,
text_encoder_fp32: bool = False,
text_encoder_mode: Optional[str] = None,
text_encoder_url: Optional[str] = None,
return_config: bool = False,
checkpoints_dir: Optional[str] = None,
):
"""Load a released Ardy model.
``modelname`` may be a short key ("core"/"g1"/"soma") or a full folder name
("Ardy-Core-RP-20FPS-Horizon40"). If a local checkpoints dir is given (via
``checkpoints_dir`` or the ``CHECKPOINTS_DIR`` env var) the model is loaded
from ``<checkpoints_dir>/<full_name>``; otherwise it is downloaded from HF.
Args:
modelname: Short key or full folder name; uses DEFAULT_MODEL if None.
device: Target device for the model (e.g. 'cuda', 'cpu').
eval_mode: If True, set model to eval mode.
default_family: Ignored (kept for call-site compatibility).
text_encoder: Pre-built text encoder to reuse, or False to load the
model without any text encoder (model.text_encoder is left as
None). When None (the default), one is built via
``load_text_encoder``.
text_encoder_fp32: If True, uses fp32 for the text encoder rather than default bfloat16.
text_encoder_mode: Backend selection ("auto"/"api"/"local"). When None,
falls back to the TEXT_ENCODER_MODE env var (default "auto").
Ignored unless ``text_encoder`` is None.
text_encoder_url: URL of the remote text-encoder service. When None,
falls back to the TEXT_ENCODER_URL env var.
checkpoints_dir: Local dir holding released model folders. When None,
falls back to the CHECKPOINTS_DIR env var; if neither is set the
model is downloaded from Hugging Face.
Returns:
Loaded model in eval mode, or (model, confg) if return_config is true
Raises:
ValueError: If modelname cannot be resolved to a released model.
FileNotFoundError: If config.yaml is missing in the model folder.
"""
if modelname is None:
modelname = DEFAULT_MODEL
# Local dir if CHECKPOINTS_DIR is set (arg or env), otherwise download from HF.
checkpoints_dir = checkpoints_dir or get_env_var("CHECKPOINTS_DIR")
# Resolve after checkpoints_dir so local-only folders (beyond the three
# released models) are accepted when loading from a local dir.
full_name = resolve_model_name(modelname, checkpoints_dir=checkpoints_dir)
if checkpoints_dir:
model_path = Path(checkpoints_dir) / full_name
if not model_path.exists():
raise FileNotFoundError(f"Model {full_name!r} not found under CHECKPOINTS_DIR {checkpoints_dir!r}.")
else:
model_path = _download_from_hf(full_name)
model_config_path = model_path / "config.yaml"
if not model_config_path.exists():
raise FileNotFoundError(f"The model folder exists but config.yaml is missing: {model_config_path}")
model_conf = OmegaConf.load(model_config_path)
# Resolve the text encoder: False means load the model without one, a
# pre-built instance is reused as-is, and None (the default) builds one
# here via load_text_encoder (which resolves a None mode through the
# TEXT_ENCODER_MODE env var). Identity checks, not truthiness: False and
# None mean different things.
if text_encoder is False:
text_encoder = None
elif text_encoder is None:
text_encoder = load_text_encoder(
mode=text_encoder_mode,
url=text_encoder_url,
fp32=text_encoder_fp32,
device=device,
)
runtime_conf = OmegaConf.create({"checkpoint_dir": str(model_path)})
model_cfg = OmegaConf.to_container(OmegaConf.merge(model_conf, runtime_conf), resolve=True)
model_cfg.pop("checkpoint_dir", None)
# The text encoder is attached after construction (or left as None for
# text_encoder=False), so prevent Hydra from instantiating one during
# construction.
model_cfg["text_encoder"] = None
model = instantiate_from_dict(model_cfg, overrides={"device": device})
if text_encoder is not None:
model.text_encoder = text_encoder
if eval_mode:
model = model.eval()
if return_config:
return model, model_cfg
return model
|