Spaces:
Sleeping
Sleeping
| """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 | |