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