"""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