Z-Image-Turbo / models.py
stanley-00's picture
Add OpenAI-compatible inference API and client
419cdc1
Raw
History Blame Contribute Delete
2.34 kB
"""Model loading and caching helpers exposed to the code environment."""
import torch
torch.set_grad_enabled(False)
def _load_text_model(model_id):
from transformers import AutoModelForCausalLM, AutoModelForSeq2SeqLM, AutoTokenizer, AutoConfig
tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
config = AutoConfig.from_pretrained(model_id, trust_remote_code=True)
arch = getattr(config, "architectures", None)
if arch and any("T5" in a or "Bart" in a or "Pegasus" in a or "Marian" in a
or "MBart" in a or "Longformer" in a for a in arch):
model_cls = AutoModelForSeq2SeqLM
else:
model_cls = AutoModelForCausalLM
model = model_cls.from_pretrained(
model_id, torch_dtype=torch.bfloat16,
low_cpu_mem_usage=True, trust_remote_code=True,
)
model.eval()
model.to("cpu")
return model, tokenizer, "cpu"
def _load_image_model(model_id):
from diffusers import DiffusionPipeline
pipe = DiffusionPipeline.from_pretrained(
model_id, torch_dtype=torch.bfloat16, low_cpu_mem_usage=True,
)
pipe.to("cpu")
if hasattr(pipe, "enable_attention_slicing"):
pipe.enable_attention_slicing()
return pipe
def _load_tts_model(model_id):
if "speecht5" in model_id.lower():
from transformers import AutoTokenizer, SpeechT5ForTextToSpeech
from datasets import load_dataset
tokenizer = AutoTokenizer.from_pretrained(model_id)
model = SpeechT5ForTextToSpeech.from_pretrained(model_id)
model.eval()
try:
emb_ds = load_dataset("Matthijs/cmu-arctic-xvectors", split="training")
speaker_embeddings = torch.tensor(emb_ds[7306]["xvector"]).unsqueeze(0)
except Exception:
speaker_embeddings = torch.zeros((1, 512))
return model, tokenizer, speaker_embeddings
else:
from transformers import pipeline as hf_pipeline
pipe = hf_pipeline("text-to-speech", model=model_id, device=-1)
return pipe, None, None
_model_cache = {}
def _cached(fn):
def wrapper(mid):
if mid not in _model_cache:
_model_cache[mid] = fn(mid)
return _model_cache[mid]
return wrapper