Spaces:
Sleeping
Sleeping
File size: 2,339 Bytes
419cdc1 | 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 | """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
|