资讯动态

AI模型打包与部署实战:从PyTorch到ONNX再到生产服务

发布时间:2026/9/11 8:41:28 来源:尧图企业网站定制
1. 这不是“上传模型就完事”的操作——AI模型管理与部署的真实战场很多人以为训练完一个模型导出个.pt或.onnx文件丢进某个框架里 run 一下就算“部署完成了”。我见过太多团队卡在这一步模型在 Jupyter Notebook 里跑得飞起一上生产环境就报CUDA out of memory、ModuleNotFoundError: No module named transformers、甚至Segmentation fault (core dumped)直接崩掉也见过业务方拿着刚训练好的文本分类模型等了三天只等到一句“还在部署中”最后发现连 Docker 镜像都没 build 成功。这不是技术不行而是把“模型管理”和“模型部署”当成了两个孤立的、一次性的运维动作。实际上它们是一条贯穿模型生命周期的、需要持续治理的流水线——从训练完成那一刻起模型就不再是“静态文件”而是一个有版本、有依赖、有性能基线、有监控指标、有回滚路径的可运行服务实体。你看到的热搜词里反复出现的ollama 部署、onnx 模型部署流程、本地部署音频转文字ai模型背后全是同一套逻辑如何让一个在实验室里诞生的模型在真实世界里稳定、高效、可控地提供价值。它不关心你用的是 ChatGPT 还是 YOLO也不区分你是跑在 Windows 11 的笔记本上还是 Kubernetes 集群里的一百台 GPU 服务器上。核心问题永远是三个这个模型能不能被准确识别和追溯它能不能在目标环境里可靠加载和执行它上线之后出了问题怎么快速定位、修复、回退这三点就是“管理和部署”的全部内涵。本文不讲抽象概念不堆砌术语只拆解我在过去三年里为金融风控、工业质检、智能客服三类场景落地的 27 个模型所踩过的坑、验证过的方案、沉淀下来的 checklist。所有内容都来自真实日志、失败截图、压测报告和凌晨三点的 Slack 沟通记录。你不需要懂 PyTorch 内部机制但必须知道torch.save()默认保存的是什么、为什么不能直接torch.load()就上线你不需要会写 Kubernetes Operator但必须明白config.toml里那几行看似无关紧要的配置如何决定你的模型服务是秒级响应还是永远 loading。我们从最基础、最容易被忽略的环节开始模型打包。2. 模型打包90% 的部署失败源于第一步就错了模型打包是模型管理的起点也是部署失败的高发区。它不是简单地把.pth文件 zip 压缩一下。真正的打包是为模型构建一个自包含、可复现、可验证的运行单元。我把它拆成四个不可跳过的层次缺一不可。2.1 层次一模型权重与结构的精确绑定PyTorch 默认的torch.save(model.state_dict(), model.pth)只保存参数不保存模型结构。这意味着你必须同时保存model.py文件并确保部署环境里的 Python 版本、PyTorch 版本、甚至__init__.py的 import 路径和训练时完全一致。这在团队协作或跨环境部署时几乎不可能保证。我吃过亏一个在torch1.13.1下训练的模型用torch2.0.1加载时nn.MultiheadAttention的内部实现变更导致forward()报错错误信息却只显示RuntimeError: expected scalar type Float but found Half根本看不出是版本问题。正确的做法是使用torch.jit.script或torch.jit.trace生成 TorchScript 模型。它将模型结构和权重一起编译成一个.pt文件脱离 Python 解释器运行。实测下来torch.jit.trace更适合有固定输入尺寸的模型如图像分类torch.jit.script更适合有控制流的模型如带 if/else 的 NLP 模型。以一个简单的文本分类模型为例# train.py import torch import torch.nn as nn class TextClassifier(nn.Module): def __init__(self, vocab_size, embed_dim, num_classes): super().__init__() self.embedding nn.Embedding(vocab_size, embed_dim) self.linear nn.Linear(embed_dim, num_classes) def forward(self, x): x self.embedding(x).mean(dim1) # 简单平均池化 return self.linear(x) # 训练完成后 model TextClassifier(vocab_size10000, embed_dim128, num_classes5) model.load_state_dict(torch.load(best_model.pth)) model.eval() # 使用 trace 打包需提供一个示例输入 example_input torch.randint(0, 10000, (1, 50)) # batch_size1, seq_len50 traced_model torch.jit.trace(model, example_input) traced_model.save(text_classifier_traced.pt)这个text_classifier_traced.pt文件就是一个独立的、可直接在任何安装了 PyTorch 的机器上加载运行的模型。它不依赖model.py也不依赖训练时的 Python 环境。这是模型可移植性的第一道保险。2.2 层次二依赖环境的精确锁定模型能跑不代表它能正确跑。一个requirements.txt文件如果只写torch而不指定torch2.0.1cu118那么在没有 CUDA 的 CPU 服务器上pip 会安装torch2.0.1的 CPU 版本而你的模型代码里可能写了model.to(cuda)结果直接RuntimeError: Found no NVIDIA driver on your system。更隐蔽的问题是transformers库的版本。transformers4.30.0和4.35.0对同一个AutoModelForSequenceClassification的forward()返回值结构可能不同导致下游推理代码崩溃。我的标准做法是在训练环境里用pip freeze requirements.lock生成一个带完整版本号和哈希值的锁定文件。这个文件不仅包含主依赖还包含所有间接依赖transitive dependencies的精确版本。例如torch2.0.1cu118 https://download.pytorch.org/whl/cu118/torch-2.0.1%2Bcu118-cp39-cp39-linux_x86_64.whl#sha256... transformers4.30.2; python_version 3.9 and python_version 4.0 tokenizers0.13.3; python_version 3.9 and python_version 4.0 ...部署时必须用pip install -r requirements.lock而不是pip install -r requirements.txt。后者是开发时的“愿望清单”前者才是生产环境的“法律合同”。2.3 层次三预处理与后处理逻辑的封装模型本身只是数学计算但一个完整的 AI 应用必然包含数据预处理如图像 resize、归一化、文本分词和后处理如 softmax 概率、NMS 非极大值抑制、标签映射。这些逻辑如果散落在 Flask API 的路由函数里或者写在 Jupyter Notebook 的 cell 里就会导致“训练-推理不一致”train-inference skew。最经典的例子训练时用transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225])推理时忘了做模型输出全是乱码。解决方案是将预处理和后处理逻辑作为模型的一部分一同打包。PyTorch 的torch.nn.Module是个强大的容器你可以把它们定义成子模块# model_with_prepost.py import torch import torch.nn as nn from torchvision import transforms from transformers import AutoTokenizer class FullPipeline(nn.Module): def __init__(self, model_path, tokenizer_namebert-base-uncased): super().__init__() # 加载训练好的模型 self.model torch.jit.load(model_path) # 封装 tokenizer这里简化实际应保存 tokenizer 的 vocab 和 config self.tokenizer AutoTokenizer.from_pretrained(tokenizer_name) # 定义预处理文本 - token ids self.preprocess lambda text: self.tokenizer( text, truncationTrue, paddingmax_length, max_length128, return_tensorspt ) # 定义后处理logits - label self.postprocess nn.Softmax(dim-1) def forward(self, texts): # 批量处理 inputs self.preprocess(texts) logits self.model(inputs[input_ids], inputs[attention_mask]) probs self.postprocess(logits) return probs # 打包整个 pipeline full_pipeline FullPipeline(text_classifier_traced.pt) full_pipeline.eval() traced_pipeline torch.jit.trace(full_pipeline, [hello world]) traced_pipeline.save(full_text_pipeline.pt)这个full_text_pipeline.pt才是一个真正开箱即用的“应用模型”。它接收原始文本字符串直接返回概率分布中间所有步骤都被固化。业务方调用时再也不用担心自己写的preprocess和训练时不一致。2.4 层次四元数据的嵌入与校验一个模型文件如果没有元数据就像一本没有书名、作者、出版日期的书。你无法回答“这个模型是哪天训练的”、“它在验证集上的准确率是多少”、“它依赖哪个版本的 CUDA”、“它的输入 shape 是什么”。这些信息必须随模型一起存储。我采用两种方式嵌入元数据文件系统层面在打包目录下强制创建metadata.json文件。内容包括{ model_name: text_classifier_v2, version: 1.2.0, train_date: 2024-05-15T14:23:00Z, accuracy_val: 0.924, input_shape: [1, 128], input_dtype: int64, output_shape: [1, 5], output_dtype: float32, hardware_requirement: {gpu: true, min_memory_gb: 4}, dependencies: {torch: 2.0.1cu118, transformers: 4.30.2} }模型文件内部对于 TorchScript 模型可以利用torch.jit.script的__dict__属性或者更稳妥地用torch.save()保存一个包含模型和元数据的字典# save_with_metadata.py model_data { model: traced_pipeline, metadata: metadata_dict, signature: sha256_of_model_weights # 可选用于完整性校验 } torch.save(model_data, text_classifier_v2_full.pt)部署脚本在加载模型前必须先读取并校验metadata.json或model_data[metadata]。例如检查hardware_requirement.gpu是否为true而当前环境没有 GPU则直接报错退出而不是等到model.to(cuda)时才崩溃。这种前置校验能将 80% 的环境不匹配问题在服务启动前就拦截掉。提示不要相信“模型文件很小随便传”的想法。一个没加元数据的.pt文件就是一个没有说明书的精密仪器。你永远不知道它需要什么条件才能运转也不知道它是否已经过时。每一次打包都是一次对模型生命周期的正式签发。3. 模型部署从“能跑”到“稳跑”的七道关卡打包解决了“模型是什么”的问题部署则要解决“模型在哪里跑、怎么跑、跑得怎么样”的问题。部署不是一键docker run而是一系列环环相扣的决策和验证。我把整个过程拆解为七个必须通过的关卡每一道关卡都对应一个真实的、高频的失败点。3.1 关卡一选择正确的推理引擎——别让框架成为瓶颈模型训练常用 PyTorch/TensorFlow但它们不是为高并发、低延迟推理设计的。直接用model.forward()在 Flask 中处理请求QPS每秒查询数可能只有个位数。你需要一个专门的推理引擎。主流选择有三个ONNX Runtime、Triton Inference Server、vLLM针对 LLM。ONNX Runtime轻量、跨平台、支持 CPU/GPU/ARM。适合中小规模、对延迟要求不极端的场景如 Web API、边缘设备。它的优势是简单把 PyTorch 模型转成 ONNX 格式然后用 ORT 加载即可。# 导出 ONNX torch.onnx.export( model, example_input, model.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}} )Triton Inference ServerNVIDIA 出品专为 GPU 推理优化支持模型并行、动态批处理dynamic batching、模型集成ensemble。适合大规模、高吞吐、多模型协同的场景如推荐系统、自动驾驶。但它需要 NVIDIA GPU 和 CUDA 环境学习成本较高。vLLM专为大语言模型LLM设计通过 PagedAttention 技术将显存利用率提升 2-4 倍显著降低推理成本。如果你部署的是Llama-2-7b、Qwen-1.5-7b这类模型vLLM 是目前事实上的标准。选择逻辑很简单看你的模型类型和业务 SLA服务等级协议。如果是一个 YOLOv8 的目标检测模型要求 API 响应时间 200msQPS 50用 ONNX Runtime TensorRTNVIDIA 的加速库是最佳组合。如果是一个 7B 参数的聊天模型要求支持 100 并发用户且希望显存占用最小vLLM 是唯一选择。我见过一个团队为了“统一技术栈”硬把 vLLM 换成 Triton结果因为 Triton 对 LLM 的 KV Cache 管理不如 vLLM显存暴涨 40%不得不回滚。3.2 关卡二容器化——Docker 不是可选项是必选项无论你用什么推理引擎都必须用 Docker 容器来封装。理由只有一个环境一致性。Windows 11 上ollama能跑不代表 Linux 服务器上也能跑本地conda env里pip install成功不代表 CI/CD 流水线里也能成功。Dockerfile 就是你的环境契约。一个健壮的 Dockerfile 必须包含基础镜像的选择优先选择官方、精简的镜像。nvidia/cuda:11.8.0-devel-ubuntu22.04比ubuntu:22.04更好因为它预装了 CUDA 工具链避免在构建时下载巨量依赖。依赖安装的分层缓存把apt-get update apt-get install -y ...和pip install -r requirements.lock分开写利用 Docker 的 layer cache加快后续构建速度。模型文件的 COPY 时机模型文件.pt,.onnx应该放在 Dockerfile 的最后几行COPY因为它们体积大、变动频繁。这样前面的依赖层可以被缓存不会因为模型更新而重新构建整个镜像。以下是一个部署 ONNX Runtime 的典型 Dockerfile# 使用 NVIDIA 官方 CUDA 基础镜像 FROM nvidia/cuda:11.8.0-devel-ubuntu22.04 # 设置环境变量 ENV DEBIAN_FRONTENDnoninteractive ENV PYTHONUNBUFFERED1 # 安装系统依赖 RUN apt-get update apt-get install -y \ python3 \ python3-pip \ python3-dev \ rm -rf /var/lib/apt/lists/* # 升级 pip 并安装 onnxruntime-gpu RUN pip3 install --upgrade pip RUN pip3 install onnxruntime-gpu1.16.0 # 创建工作目录 WORKDIR /app # 复制依赖文件会被缓存 COPY requirements.lock . RUN pip3 install -r requirements.lock # 复制应用代码 COPY app/ . # 复制模型文件最后复制避免缓存失效 COPY models/ ./models/ # 暴露端口 EXPOSE 8000 # 启动命令 CMD [python3, server.py]构建命令docker build -t text-classifier-api:v1.2.0 .。这个v1.2.0标签必须和你模型文件的version字段严格一致。这是模型版本与镜像版本对齐的关键。3.3 关卡三API 服务设计——REST 还是 gRPC不只是协议选择API 是模型与外界的唯一接口。设计不当会带来灾难性后果。最常见的错误是用 Flask 的app.route写一个/predict接口然后在request.json里直接json.loads()再喂给模型。这会导致无并发Flask 默认是单线程一个请求阻塞所有请求排队。无超时模型推理卡住API 永远不返回。无限长请求体恶意用户发送 1GB 的 JSON直接打爆内存。正确的做法是使用异步框架FastAPI 或 Starlette。它们基于asyncio能轻松处理数千并发连接。设置严格的请求体限制fastapi的Body(..., max_length1024*1024)限制 JSON body 不超过 1MB。设置推理超时在调用模型前用asyncio.wait_for()包裹超时则返回503 Service Unavailable。# server.py (FastAPI) from fastapi import FastAPI, HTTPException, Body from pydantic import BaseModel import asyncio import torch app FastAPI() class PredictRequest(BaseModel): texts: list[str] app.post(/predict) async def predict(request: PredictRequest Body(..., max_length1024*1024)): try: # 限制最大文本数 if len(request.texts) 32: raise HTTPException(status_code400, detailMax 32 texts per request) # 异步执行推理设置 5 秒超时 result await asyncio.wait_for( run_inference(request.texts), timeout5.0 ) return {result: result} except asyncio.TimeoutError: raise HTTPException(status_code503, detailInference timeout) except Exception as e: raise HTTPException(status_code500, detailstr(e)) async def run_inference(texts): # 这里调用你的 ONNX Runtime 或 TorchScript 模型 # 注意ONNX Runtime 的 Session.run() 是同步的需要用线程池包装 loop asyncio.get_event_loop() result await loop.run_in_executor(None, _sync_inference, texts) return result至于 REST 还是 gRPC取决于你的客户端。如果是 Web 前端、手机 AppREST/JSON 是标准如果是微服务内部调用如推荐服务调用 NLP 服务gRPC 的二进制协议和流式传输能带来 30% 的性能提升。3.4 关卡四资源隔离与限制——CPU/GPU 的“划片儿”一个服务器上跑多个模型服务不加资源限制就是一场灾难。一个模型因 bug 进入死循环疯狂占用 CPU其他所有服务都会变慢甚至超时。GPU 更是如此一个模型占满显存其他模型直接 OOM。Docker 的--cpus和--memory参数是基础保障docker run -d \ --cpus1.5 \ --memory2g \ --gpus device0 \ # 指定使用 GPU 0 -p 8000:8000 \ text-classifier-api:v1.2.0但这还不够。对于 GPU必须使用nvidia-docker的--gpus参数并配合NVIDIA_VISIBLE_DEVICES环境变量确保容器只能看到分配给它的 GPU 设备。更重要的是要在模型代码里显式指定device# 在推理代码中 device torch.device(cuda:0 if torch.cuda.is_available() else cpu) model model.to(device) # ... 然后所有 tensor 都要 .to(device)否则即使 Docker 限制了 GPUPyTorch 仍可能尝试访问所有可见 GPU导致冲突。3.5 关卡五健康检查与就绪探针——Kubernetes 的生命线如果你用 KubernetesK8s编排服务livenessProbe存活探针和readinessProbe就绪探针不是可选项而是 K8s 正常工作的前提。livenessProbe探测服务是否“活着”。如果失败K8s 会杀死并重启 Pod。探测/healthz返回 200 即可。readinessProbe探测服务是否“准备好”接收流量。如果失败K8s 会将该 Pod 从 Service 的 Endpoint 列表中移除不再转发请求。探测/readyz必须检查模型是否已加载、GPU 显存是否充足、依赖服务如 Redis 缓存是否连通。一个健壮的/readyz实现app.get(/readyz) def readyz(): # 1. 检查模型是否已加载 if not hasattr(app.state, model) or app.state.model is None: return {status: error, reason: model not loaded} # 2. 检查 GPU 显存如果用了 GPU if torch.cuda.is_available(): free_mem torch.cuda.mem_get_info()[0] / 1024**3 # GB if free_mem 1.0: # 小于 1GB 视为不就绪 return {status: error, reason: fGPU memory low: {free_mem:.2f} GB} # 3. 检查依赖服务如 Redis try: redis_client.ping() except Exception as e: return {status: error, reason: fredis unreachable: {e}} return {status: ok}没有readinessProbeK8s 可能在模型还没加载完时就把流量切过来导致大量 503 错误。这是线上事故的常见根源。3.6 关卡六监控与告警——看不见的故障比看得见的更可怕部署完成不等于万事大吉。模型性能会随时间漂移data drift硬件会老化网络会波动。没有监控你就是在盲人骑马。必须监控的四大黄金指标指标类别具体指标采集方式告警阈值说明可用性HTTP 5xx 错误率Prometheus nginx exporter 1% 持续 5 分钟表明服务崩溃或严重异常延迟P95 推理延迟FastAPI middleware 记录time.time() 1000ms 持续 10 分钟用户体验恶化资源GPU 显存使用率nvidia-smi Prometheus node exporter 95% 持续 15 分钟预示 OOM 风险业务模型预测置信度均值在推理代码中probs.max().item() 0.6 持续 30 分钟可能发生数据漂移我用 Grafana Prometheus 搭建监控面板所有指标都可视化。一旦GPU 显存使用率超过 95%立刻触发企业微信告警通知 SRE 团队扩容。这套系统让我们在去年一次上游数据源格式变更导致输入文本长度暴增时提前 2 小时发现了P95 延迟的缓慢爬升及时介入避免了业务中断。3.7 关卡七灰度发布与一键回滚——上线不是终点而是起点新模型上线绝不能“全量切换”。必须灰度。我的标准流程是流量切分用 Istio 或 Nginx将 1% 的流量导向新模型服务。对比监控在同一时间窗口对比新旧模型的P95 延迟、5xx 错误率、预测结果差异率对同一 batch 输入统计新旧模型输出 top-1 label 不同的比例。人工抽检随机抽取 100 条灰度流量的请求人工检查新模型的输出是否合理。逐步放大确认无问题后按 1% → 10% → 50% → 100% 的节奏放大。如果在 10% 流量时发现预测结果差异率突然从 0.1% 升到 5%立即暂停回滚。回滚不是git checkout而是kubectl rollout undo deployment/text-classifier或者docker service rollback text-classifier。整个过程必须在 2 分钟内完成。为此我要求所有模型镜像都保留至少最近 3 个版本的 tagv1.2.0,v1.1.0,v1.0.0并确保它们都能在当前集群环境中正常运行。回滚的快捷键是每个 AI 工程师的保命符。注意灰度发布不是“技术炫技”而是对业务负责。一个在测试环境完美的模型可能在生产环境遇到从未见过的脏数据、特殊字符、超长文本。灰度是唯一的、低成本的验证方式。4. 模型管理让 100 个模型像 1 个模型一样可控当团队从部署 1 个模型发展到同时维护 50 个模型不同版本、不同业务线、不同客户手动管理就成了噩梦。“那个风控模型的 v3.2 版本在哪”、“客户 A 用的是 v2.1客户 B 用的是 v3.0怎么保证不混淆”、“模型 v2.5 的训练数据集丢了怎么复现”——这些问题指向一个核心缺乏统一的模型注册中心Model Registry。4.1 为什么不能只用 Git 或文件服务器Git 适合代码不适合大型二进制模型文件。git lfs可以解决大文件但它没有模型特有的元数据管理能力如精度、输入输出 schema。文件服务器如 NFS更糟没有版本控制、没有权限管理、没有审计日志。一个误操作rm -rf就能让整个团队停工。4.2 构建自己的轻量级模型注册中心我不推荐一开始就上 MLflow 或 DVC 这类重型工具。对于中小团队一个基于 MinIO对象存储 PostgreSQL元数据库 FastAPIAPI 层的自建方案更灵活、更可控。MinIO开源的 S3 兼容对象存储。用来存.pt、.onnx等大文件。它提供版本控制、加密、生命周期管理。PostgreSQL存所有元数据。一张models表字段包括id,name,version,description,created_at,created_by,statusactive/archived一张artifacts表记录每个模型版本对应的minio_path,size,sha256_hash一张metrics表存accuracy,latency_p95,f1_score等评估指标。API 层提供标准 CRUDPOST /models上传模型自动解析metadata.json存入数据库。GET /models/{name}/versions列出所有版本。GET /models/{name}/versions/{version}获取指定版本详情和下载链接。PUT /models/{name}/versions/{version}/promote将一个版本标记为production。这个系统让模型管理变得像管理 Docker 镜像一样简单。CI/CD 流水线在模型训练完成后自动调用POST /models上传模型并记录元数据。部署脚本在拉取镜像时先调用GET /models/text-classifier/versions/latest?statusproduction拿到最新的生产版本号再用这个版本号去拉取对应的 Docker 镜像。整个链条全自动、可追溯、零人工干预。4.3 模型血缘追踪从代码到数据的完整谱系一个模型的“健康”不仅取决于它自己还取决于它的“父母”训练它的代码、它所用的数据集、它所依赖的基础镜像。这就是血缘Lineage。我的实践是在模型元数据中强制记录三个 IDcode_commit_id: 训练脚本所在的 Git commit hash。dataset_id: 数据集在 MinIO 中的唯一标识如ds-20240515-v1。base_image_id: 构建模型镜像所用的基础镜像 tag如nvidia/cuda:11.8.0-devel-ubuntu22.04。当一个模型在线上表现异常时我可以查model v1.2.0的code_commit_idgit show看那天改了什么。查dataset_id去数据湖里检查那个数据集的样本分布是否发生了偏移。查base_image_id确认是否因为基础镜像升级引入了新的 bug。这三步构成了一个完整的根因分析RCA路径。没有血缘排查就是大海捞针。4.4 模型生命周期策略不是所有模型都该永生模型不是文物不该永久保存。必须制定清晰的生命周期策略Active活跃正在生产环境使用的版本。保留所有元数据和二进制文件。Deprecated弃用已有新版本替代但旧客户仍在使用。只保留二进制文件元数据可归档。Archived归档超过 6 个月未被任何服务引用。二进制文件可压缩存储元数据保留。Deleted删除满足合规要求如 GDPR彻底清除。这个策略由一个定时任务CronJob自动执行。它每天扫描models表根据last_used_at时间戳和status字段执行相应的清理操作。自动化是管理大规模模型的唯一出路。5. 实战案例从零部署一个本地音频转文字模型Whisper现在我们把前面所有原则融合到一个具体、热门的场景中部署一个本地的、无需联网的音频转文字ASR模型。这正是热搜词本地部署音频转文字ai模型所指的需求。我们将用 OpenAI 的 Whisper 模型走一遍完整的“管理-部署”流程。5.1 场景需求分析目标用户企业内部会议记录员、听障人士辅助工具开发者。核心诉求离线、隐私安全、高准确率、支持中文。约束条件运行在 Windows 11 笔记本RTX 4060 Laptop GPU16GB RAM不依赖云服务。5.2 模型选型与打包Whisper 有多个尺寸tiny,base,small,medium,large。large最准但显存要求高10GBsmall在 RTX 40608GB 显存上更合适。我们选择openai/whisper-small。打包步骤环境准备创建whisper-envconda 环境安装transformers4.30.2,torch2.0.1cu118,ffmpeg。模型导出使用transformers的pipeline但为了可控我们手动构建模型和 processorfrom transformers import WhisperProcessor, WhisperForConditionalGeneration import torch processor WhisperProcessor.from_pretrained(openai/whisper-small, languagechinese, tasktranscribe) model WhisperForConditionalGeneration.from_pretrained(openai/whisper-small) # 转换为 TorchScripttrace # Whisper 的输入是 mel spectrogram需要构造一个 dummy input dummy_mel torch.randn(1, 80, 3000) # batch1, n_mels80, time_steps3000 dummy_input_ids torch.tensor([[1]]) # decoder start token traced_model torch.jit.trace(model, (dummy_mel, dummy_input_ids)) traced_model.save(whisper_small_traced.pt)封装 Pipeline创建WhisperPipeline类整合processor的feature_extractor和tokenizer以及traced_model的generate逻辑。生成元数据metadata.json包含model_name: whisper-small-zh,input_format: wav, max

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价