返回模型
说明文档
该模型是 Pygmalion-6b 的 ONNX 导出版本,所有功劳应归于 PygmalionAI。
请注意,此 ONNX 导出版本并非完全精确,由于 PyTorch ONNX 导出的限制,它被上采样为 Float32,这将占用原始 Pygmalion AI 模型两倍的内存。此导出的目的是获取算子和节点列表,然后可用于在 Vulkan Compute 上对 Pygmalion 6b 模型运行推理,最终实现无需繁琐的推理,同时支持 INT8 或 INT4 量化,并兼容几乎所有支持 Vulkan Compute 的设备。
以下是相关脚本,model.py 获取自 PygmalioniAI/gradio-ui,并根据 GNU Affero General Public License v3.0 许可。出于对该许可的尊重,下面列出的所有脚本均遵循 GNU Affero General Public License v3.0。
export.py
import torch
import onnx
import transformers
import typing as t
model_name = \"PygmalionAI/pygmalion-6b\"
from model import build_model_and_tokenizer_for, run_raw_inference
model, tokenizer = build_model_and_tokenizer_for(model_name)
model.to('cpu').float()
input_layer = model.get_input_embeddings()
output_layer = model.get_output_embeddings()
# Load PyTorch model from .pth file
#model = AutoModelForCausalLM.from_pretrained(\"PygmalionAI/pygmalion-6b\")
#state_dict = torch.load('pygmalion-6b.pth')
#model.load_state_dict(state_dict)
# Export PyTorch model to ONNX format
# Encode some input text
input_text = \"Hello, how are you today?\"
encoded_input = tokenizer.encode(input_text, return_tensors='pt')
# Export the tokenizer to ONNX format
print(f\"Raw: {input_text}\")
print(f\"Encoded: {encoded_input}\")
output_path = \"onnx/pygmalion-6b.onnx\"
dummy_input = torch.zeros((1, 10), dtype=torch.long)
input_names = [\"input_ids\"]
output_names = [\"output\"]
dynamic_axes = {\"input_ids\": {0: \"batch_size\", 1: \"sequence_length\"},
\"output\": {0: \"batch_size\", 1: \"sequence_length\"}}
torch.onnx.export(model, dummy_input, output_path, input_names=input_names,
output_names=output_names, dynamic_axes=dynamic_axes,
opset_version=12)
model.py
import logging
import typing as t
import torch
import transformers
logger = logging.getLogger(__name__)
def build_model_and_tokenizer_for(
model_name: str
) -> t.Tuple[transformers.AutoModelForCausalLM, transformers.AutoTokenizer]:
'''Sets up the model and accompanying objects.'''
logger.info(f\"Loading tokenizer for {model_name}\")
tokenizer = transformers.AutoTokenizer.from_pretrained(model_name)
# NOTE(11b): non-OPT models support passing this in at inference time, might
# be worth refactoring for a debug version so we're able to experiment on
# the fly
bad_words_ids = [
tokenizer(bad_word, add_special_tokens=False).input_ids
for bad_word in _build_bad_words_list_for(model_name)
]
logger.info(f\"Loading the {model_name} model\")
model = transformers.AutoModelForCausalLM.from_pretrained(
model_name, bad_words_ids=bad_words_ids)
model.eval().to(\"cpu\")
logger.info(\"Model and tokenizer are ready\")
return model, tokenizer
def build_tokenizer_for(
model_name: str
) -> t.Tuple[transformers.AutoTokenizer]:
'''Sets up the model and accompanying objects.'''
logger.info(f\"Loading tokenizer for {model_name}\")
tokenizer = transformers.AutoTokenizer.from_pretrained(model_name)
# NOTE(11b): non-OPT models support passing this in at inference time, might
# be worth refactoring for a debug version so we're able to experiment on
# the fly
bad_words_ids = [
tokenizer(bad_word, add_special_tokens=False).input_ids
for bad_word in _build_bad_words_list_for(model_name)
]
return tokenizer
def run_raw_inference(model: transformers.AutoModelForCausalLM,
tokenizer: transformers.AutoTokenizer, prompt: str,
user_message: str, **kwargs: t.Any) -> str:
'''
Runs inference on the model, and attempts to returns only the newly
generated text.
:param model: Model to perform inference with.
:param tokenizer: Tokenizer to tokenize input with.
:param prompt: Input to feed to the model.
:param user_message: The user's raw message, exactly as appended to the end
of `prompt`. Used for trimming the original input from the model output.
:return: Decoded model generation.
'''
tokenized_items = tokenizer(prompt, return_tensors=\"pt\").to(\"cpu\")
# Atrocious code to stop generation when the model outputs \"\nYou: \" in
# freshly generated text. Feel free to send in a PR if you know of a
# cleaner way to do this.
stopping_criteria_list = transformers.StoppingCriteriaList([
_SentinelTokenStoppingCriteria(
sentinel_token_ids=tokenizer(
\"\nYou:\",
add_special_tokens=False,
return_tensors=\"pt\",
).input_ids.to(\"cpu\"),
starting_idx=tokenized_items.input_ids.shape[-1])
])
logits = model.generate(stopping_criteria=stopping_criteria_list,
**tokenized_items,
**kwargs)
output = tokenizer.decode(logits[0], skip_special_tokens=True)
logger.debug(\"Before trimming, model output was: `%s`\", output)
# Trim out the input prompt from the generated output.
if (idx := prompt.rfind(user_message)) != -1:
trimmed_output = output[idx + len(user_message) - 1:].strip()
logger.debug(\"After trimming, it became: `%s`\", trimmed_output)
return trimmed_output
else:
raise Exception(
\"Couldn't find user message in the model's output. What?\")
def _build_bad_words_list_for(_model_name: str) -> t.List[str]:
'''Builds a list of bad words for the given model.'''
# NOTE(11b): This was implemented as a function because each model size
# seems to have it quirks at the moment, but this is a rushed implementation
# so I'm not handling that, hence the dumb return here.
return [\"Persona:\", \"Scenario:\", \"<START>\"]
#class _SentinelTokenStoppingCriteria(transformers.StoppingCriteria):
# def __init__(self, sentinel_token_ids: torch.LongTensor,
# starting_idx: int):
# transformers.StoppingCriteria.__init__(self)
# self.sentinel_token_ids = sentinel_token_ids
# self.starting_idx = starting_idx
# def __call__(self, input_ids: torch.LongTensor,
# _scores: torch.FloatTensor) -> bool:
# for sample in input_ids:
# trimmed_sample = sample[self.starting_idx:]
# # Can't unfold, output is still too tiny. Skip.
# if trimmed_sample.shape[-1] < self.sentinel_token_ids.shape[-1]:
# continue
# for window in trimmed_sample.unfold(
# 0, self.sentinel_token_ids.shape[-1], 1):
# if torch.all(torch.eq(self.sentinel_token_ids, window)):
# return True
# return False
BCenti/pygmalion_onnx_export
作者 BCenti
↓ 0
♥ 0
创建时间: 2023-03-19 02:04:56+00:00
更新时间: 2023-03-19 07:06:27+00:00
在 Hugging Face 上查看文件 (287)
.gitattributes
README.md
h.0.attn.bias
h.0.ln_1.bias
h.0.ln_1.weight
h.0.mlp.fc_in.bias
h.0.mlp.fc_out.bias
h.1.ln_1.bias
h.1.ln_1.weight
h.1.mlp.fc_in.bias
h.1.mlp.fc_out.bias
h.10.ln_1.bias
h.10.ln_1.weight
h.10.mlp.fc_in.bias
h.10.mlp.fc_out.bias
h.11.ln_1.bias
h.11.ln_1.weight
h.11.mlp.fc_in.bias
h.11.mlp.fc_out.bias
h.12.ln_1.bias
h.12.ln_1.weight
h.12.mlp.fc_in.bias
h.12.mlp.fc_out.bias
h.13.ln_1.bias
h.13.ln_1.weight
h.13.mlp.fc_in.bias
h.13.mlp.fc_out.bias
h.14.ln_1.bias
h.14.ln_1.weight
h.14.mlp.fc_in.bias
h.14.mlp.fc_out.bias
h.15.ln_1.bias
h.15.ln_1.weight
h.15.mlp.fc_in.bias
h.15.mlp.fc_out.bias
h.16.ln_1.bias
h.16.ln_1.weight
h.16.mlp.fc_in.bias
h.16.mlp.fc_out.bias
h.17.ln_1.bias
h.17.ln_1.weight
h.17.mlp.fc_in.bias
h.17.mlp.fc_out.bias
h.18.ln_1.bias
h.18.ln_1.weight
h.18.mlp.fc_in.bias
h.18.mlp.fc_out.bias
h.19.ln_1.bias
h.19.ln_1.weight
h.19.mlp.fc_in.bias
h.19.mlp.fc_out.bias
h.2.ln_1.bias
h.2.ln_1.weight
h.2.mlp.fc_in.bias
h.2.mlp.fc_out.bias
h.20.ln_1.bias
h.20.ln_1.weight
h.20.mlp.fc_in.bias
h.20.mlp.fc_out.bias
h.21.ln_1.bias
h.21.ln_1.weight
h.21.mlp.fc_in.bias
h.21.mlp.fc_out.bias
h.22.ln_1.bias
h.22.ln_1.weight
h.22.mlp.fc_in.bias
h.22.mlp.fc_out.bias
h.23.ln_1.bias
h.23.ln_1.weight
h.23.mlp.fc_in.bias
h.23.mlp.fc_out.bias
h.24.ln_1.bias
h.24.ln_1.weight
h.24.mlp.fc_in.bias
h.24.mlp.fc_out.bias
h.25.ln_1.bias
h.25.ln_1.weight
h.25.mlp.fc_in.bias
h.25.mlp.fc_out.bias
h.26.ln_1.bias
h.26.ln_1.weight
h.26.mlp.fc_in.bias
h.26.mlp.fc_out.bias
h.27.ln_1.bias
h.27.ln_1.weight
h.27.mlp.fc_in.bias
h.27.mlp.fc_out.bias
h.3.ln_1.bias
h.3.ln_1.weight
h.3.mlp.fc_in.bias
h.3.mlp.fc_out.bias
h.4.ln_1.bias
h.4.ln_1.weight
h.4.mlp.fc_in.bias
h.4.mlp.fc_out.bias
h.5.ln_1.bias
h.5.ln_1.weight
h.5.mlp.fc_in.bias
h.5.mlp.fc_out.bias
h.6.ln_1.bias
h.6.ln_1.weight
h.6.mlp.fc_in.bias
h.6.mlp.fc_out.bias
h.7.ln_1.bias
h.7.ln_1.weight
h.7.mlp.fc_in.bias
h.7.mlp.fc_out.bias
h.8.ln_1.bias
h.8.ln_1.weight
h.8.mlp.fc_in.bias
h.8.mlp.fc_out.bias
h.9.ln_1.bias
h.9.ln_1.weight
h.9.mlp.fc_in.bias
h.9.mlp.fc_out.bias
ln_f.bias
ln_f.weight
model.onnx
ONNX
onnx__MatMul_10003
onnx__MatMul_10004
onnx__MatMul_10005
onnx__MatMul_10006
onnx__MatMul_10007
onnx__MatMul_10008
onnx__MatMul_10033
onnx__MatMul_10034
onnx__MatMul_10035
onnx__MatMul_10036
onnx__MatMul_10037
onnx__MatMul_10038
onnx__MatMul_10063
onnx__MatMul_10064
onnx__MatMul_10065
onnx__MatMul_10066
onnx__MatMul_10067
onnx__MatMul_10068
onnx__MatMul_10093
onnx__MatMul_10094
onnx__MatMul_10095
onnx__MatMul_10096
onnx__MatMul_10097
onnx__MatMul_10098
onnx__MatMul_10123
onnx__MatMul_10124
onnx__MatMul_10125
onnx__MatMul_10126
onnx__MatMul_10127
onnx__MatMul_10128
onnx__MatMul_10153
onnx__MatMul_10154
onnx__MatMul_10155
onnx__MatMul_10156
onnx__MatMul_10157
onnx__MatMul_10158
onnx__MatMul_10183
onnx__MatMul_10184
onnx__MatMul_10185
onnx__MatMul_10186
onnx__MatMul_10187
onnx__MatMul_10188
onnx__MatMul_10213
onnx__MatMul_10214
onnx__MatMul_10215
onnx__MatMul_10216
onnx__MatMul_10217
onnx__MatMul_10218
onnx__MatMul_10243
onnx__MatMul_10244
onnx__MatMul_10245
onnx__MatMul_10246
onnx__MatMul_10247
onnx__MatMul_10248
onnx__MatMul_10273
onnx__MatMul_10274
onnx__MatMul_10275
onnx__MatMul_10276
onnx__MatMul_10277
onnx__MatMul_10278
onnx__MatMul_10303
onnx__MatMul_10304
onnx__MatMul_10305
onnx__MatMul_10306
onnx__MatMul_10307
onnx__MatMul_10308
onnx__MatMul_10333
onnx__MatMul_10334
onnx__MatMul_10335
onnx__MatMul_10336
onnx__MatMul_10337
onnx__MatMul_10338
onnx__MatMul_10363
onnx__MatMul_10364
onnx__MatMul_10365
onnx__MatMul_10366
onnx__MatMul_10367
onnx__MatMul_10368
onnx__MatMul_10393
onnx__MatMul_10394
onnx__MatMul_10395
onnx__MatMul_10396
onnx__MatMul_10397
onnx__MatMul_10398
onnx__MatMul_10423
onnx__MatMul_10424
onnx__MatMul_10425
onnx__MatMul_10426
onnx__MatMul_10427
onnx__MatMul_10428
onnx__MatMul_10453
onnx__MatMul_10454
onnx__MatMul_10455
onnx__MatMul_10456
onnx__MatMul_10457
onnx__MatMul_10458
onnx__MatMul_10483
onnx__MatMul_10484
onnx__MatMul_10485
onnx__MatMul_10486
onnx__MatMul_10487
onnx__MatMul_10488
onnx__MatMul_10513
onnx__MatMul_10514
onnx__MatMul_10515
onnx__MatMul_10516
onnx__MatMul_10517
onnx__MatMul_10518
onnx__MatMul_10543
onnx__MatMul_10544
onnx__MatMul_10545
onnx__MatMul_9706
onnx__MatMul_9707
onnx__MatMul_9708
onnx__MatMul_9733
onnx__MatMul_9734
onnx__MatMul_9735
onnx__MatMul_9736
onnx__MatMul_9737
onnx__MatMul_9738
onnx__MatMul_9763
onnx__MatMul_9764
onnx__MatMul_9765
onnx__MatMul_9766
onnx__MatMul_9767
onnx__MatMul_9768
onnx__MatMul_9793
onnx__MatMul_9794
onnx__MatMul_9795
onnx__MatMul_9796
onnx__MatMul_9797
onnx__MatMul_9798
onnx__MatMul_9823
onnx__MatMul_9824
onnx__MatMul_9825
onnx__MatMul_9826
onnx__MatMul_9827
onnx__MatMul_9828
onnx__MatMul_9853
onnx__MatMul_9854
onnx__MatMul_9855
onnx__MatMul_9856
onnx__MatMul_9857
onnx__MatMul_9858
onnx__MatMul_9883
onnx__MatMul_9884
onnx__MatMul_9885
onnx__MatMul_9886
onnx__MatMul_9887
onnx__MatMul_9888
onnx__MatMul_9913
onnx__MatMul_9914
onnx__MatMul_9915
onnx__MatMul_9916
onnx__MatMul_9917
onnx__MatMul_9918
onnx__MatMul_9943
onnx__MatMul_9944
onnx__MatMul_9945
onnx__MatMul_9946
onnx__MatMul_9947
onnx__MatMul_9948
onnx__MatMul_9973
onnx__MatMul_9974
onnx__MatMul_9975
onnx__MatMul_9976
onnx__MatMul_9977
onnx__MatMul_9978
wte.weight