ONNX 模型库
返回模型

说明文档

该模型是 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