Views
No views yet
1import torch
2import onnx
3import transformers
4import typing as t
5
6model_name = "PygmalionAI/pygmalion-6b"
7from model import build_model_and_tokenizer_for, run_raw_inference
8model, tokenizer = build_model_and_tokenizer_for(model_name)
9model.to('cpu').float()
10
11input_layer = model.get_input_embeddings()
12output_layer = model.get_output_embeddings()
13
14# Load PyTorch model from .pth file
15#model = AutoModelForCausalLM.from_pretrained("PygmalionAI/pygmalion-6b")
16
17#state_dict = torch.load('pygmalion-6b.pth')
18
19#model.load_state_dict(state_dict)
20
21# Export PyTorch model to ONNX format
22# Encode some input text
23input_text = "Hello, how are you today?"
24encoded_input = tokenizer.encode(input_text, return_tensors='pt')
25
26# Export the tokenizer to ONNX format
27print(f"Raw: {input_text}")
28print(f"Encoded: {encoded_input}")
29
30output_path = "onnx/pygmalion-6b.onnx"
31dummy_input = torch.zeros((1, 10), dtype=torch.long)
32input_names = ["input_ids"]
33output_names = ["output"]
34dynamic_axes = {"input_ids": {0: "batch_size", 1: "sequence_length"},
35 "output": {0: "batch_size", 1: "sequence_length"}}
36torch.onnx.export(model, dummy_input, output_path, input_names=input_names,
37 output_names=output_names, dynamic_axes=dynamic_axes,
38 opset_version=12)1import logging
2import typing as t
3
4import torch
5import transformers
6
7logger = logging.getLogger(__name__)
8
9
10def build_model_and_tokenizer_for(
11 model_name: str
12) -> t.Tuple[transformers.AutoModelForCausalLM, transformers.AutoTokenizer]:
13 '''Sets up the model and accompanying objects.'''
14 logger.info(f"Loading tokenizer for {model_name}")
15 tokenizer = transformers.AutoTokenizer.from_pretrained(model_name)
16
17 # NOTE(11b): non-OPT models support passing this in at inference time, might
18 # be worth refactoring for a debug version so we're able to experiment on
19 # the fly
20 bad_words_ids = [
21 tokenizer(bad_word, add_special_tokens=False).input_ids
22 for bad_word in _build_bad_words_list_for(model_name)
23 ]
24
25 logger.info(f"Loading the {model_name} model")
26 model = transformers.AutoModelForCausalLM.from_pretrained(
27 model_name, bad_words_ids=bad_words_ids)
28 model.eval().to("cpu")
29
30 logger.info("Model and tokenizer are ready")
31 return model, tokenizer
32
33def build_tokenizer_for(
34 model_name: str
35) -> t.Tuple[transformers.AutoTokenizer]:
36 '''Sets up the model and accompanying objects.'''
37 logger.info(f"Loading tokenizer for {model_name}")
38 tokenizer = transformers.AutoTokenizer.from_pretrained(model_name)
39
40 # NOTE(11b): non-OPT models support passing this in at inference time, might
41 # be worth refactoring for a debug version so we're able to experiment on
42 # the fly
43 bad_words_ids = [
44 tokenizer(bad_word, add_special_tokens=False).input_ids
45 for bad_word in _build_bad_words_list_for(model_name)
46 ]
47
48 return tokenizer
49
50
51def run_raw_inference(model: transformers.AutoModelForCausalLM,
52 tokenizer: transformers.AutoTokenizer, prompt: str,
53 user_message: str, **kwargs: t.Any) -> str:
54 '''
55 Runs inference on the model, and attempts to returns only the newly
56 generated text.
57
58 :param model: Model to perform inference with.
59 :param tokenizer: Tokenizer to tokenize input with.
60 :param prompt: Input to feed to the model.
61 :param user_message: The user's raw message, exactly as appended to the end
62 of `prompt`. Used for trimming the original input from the model output.
63 :return: Decoded model generation.
64 '''
65 tokenized_items = tokenizer(prompt, return_tensors="pt").to("cpu")
66
67 # Atrocious code to stop generation when the model outputs "\nYou: " in
68 # freshly generated text. Feel free to send in a PR if you know of a
69 # cleaner way to do this.
70 stopping_criteria_list = transformers.StoppingCriteriaList([
71 _SentinelTokenStoppingCriteria(
72 sentinel_token_ids=tokenizer(
73 "\nYou:",
74 add_special_tokens=False,
75 return_tensors="pt",
76 ).input_ids.to("cpu"),
77 starting_idx=tokenized_items.input_ids.shape[-1])
78 ])
79
80 logits = model.generate(stopping_criteria=stopping_criteria_list,
81 **tokenized_items,
82 **kwargs)
83 output = tokenizer.decode(logits[0], skip_special_tokens=True)
84
85 logger.debug("Before trimming, model output was: `%s`", output)
86
87 # Trim out the input prompt from the generated output.
88 if (idx := prompt.rfind(user_message)) != -1:
89 trimmed_output = output[idx + len(user_message) - 1:].strip()
90 logger.debug("After trimming, it became: `%s`", trimmed_output)
91
92 return trimmed_output
93 else:
94 raise Exception(
95 "Couldn't find user message in the model's output. What?")
96
97
98def _build_bad_words_list_for(_model_name: str) -> t.List[str]:
99 '''Builds a list of bad words for the given model.'''
100
101 # NOTE(11b): This was implemented as a function because each model size
102 # seems to have it quirks at the moment, but this is a rushed implementation
103 # so I'm not handling that, hence the dumb return here.
104 return ["Persona:", "Scenario:", "<START>"]
105
106
107#class _SentinelTokenStoppingCriteria(transformers.StoppingCriteria):
108
109# def __init__(self, sentinel_token_ids: torch.LongTensor,
110# starting_idx: int):
111# transformers.StoppingCriteria.__init__(self)
112# self.sentinel_token_ids = sentinel_token_ids
113# self.starting_idx = starting_idx
114
115# def __call__(self, input_ids: torch.LongTensor,
116# _scores: torch.FloatTensor) -> bool:
117# for sample in input_ids:
118# trimmed_sample = sample[self.starting_idx:]
119# # Can't unfold, output is still too tiny. Skip.
120# if trimmed_sample.shape[-1] < self.sentinel_token_ids.shape[-1]:
121# continue
122
123# for window in trimmed_sample.unfold(
124# 0, self.sentinel_token_ids.shape[-1], 1):
125# if torch.all(torch.eq(self.sentinel_token_ids, window)):
126# return True
127# return False