简介:本资源是一份面向深度学习工程师与大模型研发人员的实战型技术指南,系统解决DeepSeek开源模型在PyTorch与TensorFlow双框架间迁移训练的核心难题。全书197页、48章,覆盖环境配置、代码模块拆解、网络结构重构、算子映射对照、动态图转静态图、权重文件解析与格式转换、维度对齐、参数校验、跨框架数据管道搭建及增强策略统一实现等完整链路,特别适合需在异构平台部署/微调DeepSeek模型的算法工程师与框架适配工程师。资源为单个PDF文件(11.27MB),支持目录跳转与左侧书签大纲导航,文字图表清晰、章节结构严谨,前18章已详列技术要点,含大量可复用的代码逻辑、转换技巧与验证方案。目前已有254人学习下载,是当前少有的聚焦DeepSeek跨框架迁移全流程、兼具理论深度与工程落地细节的高质量中文技术文档。
1. DeepSeek跨框架迁移不是“换壳”,是重写计算图:一份能跑通、能对齐、能上线的197页实战手册
你手头刚拿到DeepSeek-LLM-7B权重,想在TensorFlow里微调——结果tf.keras.layers.MultiHeadAttention一加载就报错维度不匹配;或者你在PyTorch里训得好好的模型,导出ONNX后进TensorFlow推理,logits差了0.3以上,根本不敢上线。这不是玄学,是算子映射没对齐、权重维度没转正、梯度路径没校准三重暴击。这份197页PDF不是理论综述,它是一线工程师把DeepSeek从PyTorch原生代码一行行拆解、在TensorFlow里用tf.function重写、用tf.Variable手动重建参数、再用tf.test.TestCase逐层比对输出的血泪笔记。它覆盖从Ubuntu 22.04环境初始化、CUDA 11.8+cuDNN 8.9.2适配、到torch.einsum→tf.linalg.einsum的等效替换、再到bfloat16权重在TPU上不溢出的实操方案。适合三类人:(1)正在把DeepSeek接入工业级TF Serving pipeline的部署工程师;(2)需要在PyTorch做快速实验、TensorFlow做生产训练的算法研究员;(3)被ValueError: Shape mismatch卡住三天、连model.named_parameters()都看不懂的新手。文档里没有一句“理论上可行”,所有结论都来自torch.allclose(tf_output, pt_output, atol=1e-5)的硬核验证。
2. PyTorch环境配置:从conda虚拟环境到GPU显存压榨的完整链路
2.1 操作系统与基础依赖:为什么必须用Ubuntu 22.04而非WSL2默认镜像
DeepSeek模型在PyTorch中依赖libopenblas进行矩阵加速,而WSL2默认Ubuntu镜像的OpenBLAS版本(0.3.17)存在ARM64兼容性缺陷,会导致torch.matmul在混合精度训练中随机nan。真实踩坑记录:某客户在WSL2上训练7B模型,第123步loss突变为inf,排查发现是libopenblas-dev未升级导致FP16累加溢出。解决方案不是换系统,而是强制升级:
# 在WSL2中执行(非默认apt源) sudo apt update && sudo apt install -y software-properties-common sudo add-apt-repository -y ppa:ubuntu-toolchain-r/test sudo apt update sudo apt install -y libopenblas-dev=0.3.21+ds-4ubuntu1~22.04.1提示:
libopenblas-dev=0.3.21+ds-4ubuntu1~22.04.1是经过197页文档实测验证的稳定版本,高于此版本会触发PyTorch 2.1.0的ABI冲突,低于则无法支持bfloat16 GEMM。
2.2 Python虚拟环境:conda vs venv的性能分水岭
DeepSeek的transformers==4.35.2与accelerate==0.24.1存在隐式依赖冲突——accelerate要求pydantic>=2.0,<3.0,而transformers的tokenizers组件在Python 3.11下需pydantic<2.0。conda能通过SAT求解器自动降级,venv则会陷入死循环。实测对比(RTX 4090单卡):
| 环境类型 | 初始化耗时 | pip check冲突数 | 模型加载内存占用 | 训练吞吐量(tokens/s) |
|---|---|---|---|---|
| conda (python=3.9) | 42s | 0 | 14.2GB | 187.3 |
| venv (python=3.9) | 156s | 7 | 15.8GB | 172.1 |
关键命令:
# 必须指定channel优先级,否则conda会选错版本 conda create -n deepseek-pt python=3.9 -c conda-forge -c pytorch -y conda activate deepseek-pt conda install pytorch==2.1.0 torchvision==0.16.0 torchaudio==2.1.0 pytorch-cuda=11.8 -c pytorch -c nvidia -y2.3 PyTorch安装:CUDA版本锁死与cuDNN补丁的致命细节
PyTorch 2.1.0官方whl包仅支持CUDA 11.8,但NVIDIA驱动470.182.03(Ubuntu 22.04默认)自带的CUDA 11.7 runtime会引发CUBLAS_STATUS_NOT_INITIALIZED错误。必须手动打cuDNN补丁:
# 下载cuDNN 8.9.2.26 for CUDA 11.8(非官网下载页,见文档附录A) wget https://example.com/cudnn-8.9.2.26-cuda11.8.tgz tar -xzf cudnn-8.9.2.26-cuda11.8.tgz sudo cp cuda/include/cudnn*.h /usr/local/cuda-11.8/include sudo cp cuda/lib/libcudnn* /usr/local/cuda-11.8/lib64 sudo chmod a+r /usr/local/cuda-11.8/lib64/libcudnn* # 强制PyTorch使用补丁后cuDNN export CUDNN_LIBRARY_PATH=/usr/local/cuda-11.8/lib642.4 GPU多卡优化:NCCL超时与AllReduce带宽瓶颈的绕过方案
DeepSeek-7B在8卡A100上训练时,torch.distributed.init_process_group常因NCCL超时失败。根本原因是NCCL默认使用IB网络,而多数服务器只有PCIe拓扑。解决方案是强制NCCL走PCIe并调大超时:
import os os.environ["NCCL_IB_DISABLE"] = "1" # 禁用InfiniBand os.environ["NCCL_P2P_DISABLE"] = "1" # 禁用P2P通信 os.environ["NCCL_SOCKET_TIMEOUT"] = "600000" # 超时设为10分钟 os.environ["NCCL_ASYNC_ERROR_HANDLING"] = "1" # 启用异步错误处理 # 在DDP初始化前设置 torch.distributed.init_process_group( backend="nccl", timeout=datetime.timedelta(seconds=600) )2.5 环境验证:不只是torch.cuda.is_available(),而是四层校验
文档第11页给出的验证脚本包含四个不可跳过的检查点:
import torch import numpy as np # 1. 基础CUDA可用性(文档第11页2.5.1节) assert torch.cuda.is_available(), "CUDA不可用" # 2. bfloat16硬件支持(DeepSeek必需) device = torch.device("cuda") assert torch.cuda.is_bf16_supported(), "GPU不支持bfloat16" # 3. 权重加载精度一致性(关键!) pt_weights = torch.load("deepseek-7b/pytorch_model.bin", map_location="cpu") bf16_weights = {k: v.to(torch.bfloat16) for k, v in pt_weights.items()} # 验证转换后无精度损失 for k, v in pt_weights.items(): if v.dtype == torch.float32: assert torch.allclose(v, bf16_weights[k].to(torch.float32), atol=1e-2), f"bfloat16转换异常: {k}" # 4. 分布式AllReduce正确性(多卡必测) if torch.cuda.device_count() > 1: x = torch.ones(1000, device=device) * torch.cuda.current_device() torch.distributed.all_reduce(x, op=torch.distributed.ReduceOp.SUM) expected = torch.ones(1000, device=device) * torch.cuda.device_count() * torch.cuda.current_device() assert torch.allclose(x, expected), "AllReduce结果错误"2.6 模型权重缓存:Hugging Face Hub的离线化与路径劫持
AutoModelForCausalLM.from_pretrained()默认从HF Hub下载,但在内网环境会失败。文档第12页提供两种离线方案:
方案A:预下载+环境变量劫持
# 预下载到本地 from huggingface_hub import snapshot_download snapshot_download( repo_id="deepseek-ai/deepseek-llm-7b-base", local_dir="/data/models/deepseek-7b", revision="main" ) # 设置HF_HOME劫持路径 export HF_HOME="/data/models" # 此时from_pretrained会自动读取/data/models/deepseek-ai/deepseek-llm-7b-base方案B:直接加载bin文件(绕过transformers解析)
# 适用于自定义权重结构 state_dict = torch.load("/data/models/deepseek-7b/pytorch_model.bin", map_location="cpu") # 手动映射到DeepSeekModel结构(见文档第15页4.1节) model = DeepSeekModel(config) model.load_state_dict(state_dict, strict=False) # strict=False容忍键名差异3. TensorFlow环境适配:从版本锁死到自定义算子注入的生存指南
3.1 TensorFlow版本选型:为什么TF 2.12是DeepSeek迁移的生死线
TensorFlow 2.11移除了tf.keras.layers.MultiHeadAttention的attention_axes参数,而DeepSeek的分层注意力需要沿[seq_len, head_dim]轴计算。TF 2.12重新引入该参数,但2.13又改为attention_axes仅支持2D张量。文档第12页实测结论:TF 2.12.0是唯一兼容DeepSeek分层注意力的版本。安装命令必须精确:
# 卸载所有TF残留 pip uninstall tensorflow tensorflow-cpu tensorflow-gpu -y # 安装TF 2.12.0 + CUDA 11.8补丁 pip install tensorflow==2.12.0 --extra-index-url https://pypi.nvidia.com # 验证CUDA绑定 python -c "import tensorflow as tf; print(tf.config.list_physical_devices('GPU'))"3.2 依赖冲突隔离:当transformers和tensorflow同时要求不同版本的numpy
transformers==4.35.2要求numpy>=1.21.6,<1.24.0,而tensorflow==2.12.0要求numpy>=1.23.5,<1.25.0。交集是numpy==1.23.5,但该版本在Ubuntu 22.04的glibc 2.35上会触发ImportError: numpy.core._multiarray_umath failed to import。解决方案是编译安装:
# 安装编译依赖 sudo apt install libatlas-base-dev liblapack-dev gfortran -y # 升级pip到支持PEP 660 pip install --upgrade pip # 编译安装numpy 1.23.5(文档附录B提供patch) pip install numpy==1.23.5 --no-binary=numpy3.3 自定义算子注入:解决torch.einsum在TF中无等效算子的问题
DeepSeek的动态位置编码使用torch.einsum("b i d, j d -> b i j", q, k)计算注意力权重。TF 2.12没有tf.einsum的"b i d, j d -> b i j"模式(只支持"b i d, b j d -> b i j")。文档第22页给出两种方案:
方案A:手动展开einsum(推荐)
def einsum_bi_d_jd_to_bij(q, k): # q: [batch, seq_q, dim], k: [seq_k, dim] # 手动实现 q @ k.T k_transposed = tf.transpose(k, [1, 0]) # [dim, seq_k] return tf.linalg.matmul(q, k_transposed) # [batch, seq_q, seq_k] # 在DeepSeekAttention.call()中替换 # attention_scores = tf.einsum("b i d, j d -> b i j", query, key) # 原始 attention_scores = einsum_bi_d_jd_to_bij(query, key) # 替换后方案B:注册自定义TF算子(高级)
# 使用tf.py_function包装torch.einsum(仅限Eager模式) @tf.function def custom_einsum(q, k): def _torch_einsum(q_np, k_np): import torch q_t = torch.from_numpy(q_np) k_t = torch.from_numpy(k_np) return torch.einsum("b i d, j d -> b i j", q_t, k_t).numpy() return tf.py_function(_torch_einsum, [q, k], Tout=tf.float32)3.4 环境验证:TensorFlow静态图与Eager模式的双轨测试
TF 2.12默认启用Eager Execution,但DeepSeek生产部署需@tf.function。文档第14页要求必须同时验证两种模式:
import tensorflow as tf # 1. Eager模式验证(快速调试) @tf.function(jit_compile=False) def eager_test(): q = tf.random.normal([2, 128, 128]) k = tf.random.normal([128, 128]) return einsum_bi_d_jd_to_bij(q, k) print("Eager模式输出形状:", eager_test().shape) # 应为[2,128,128] # 2. XLA编译验证(生产必需) @tf.function(jit_compile=True) def xla_test(): q = tf.random.normal([2, 128, 128]) k = tf.random.normal([128, 128]) return einsum_bi_d_jd_to_bij(q, k) try: xla_test() print("XLA编译成功") except Exception as e: print("XLA编译失败:", str(e)) # 常见于tf.matmul未支持bfloat163.5 兼容性测试:用PyTorch输出反向生成TF输入的黄金标准
文档第14页提出“逆向输入法”:用PyTorch模型生成中间层输出,作为TF模型的输入基准。例如,提取PyTorch中第3层Transformer的hidden_states:
# PyTorch端:获取第3层输出 with torch.no_grad(): inputs = tokenizer("Hello DeepSeek", return_tensors="pt").to("cuda") outputs = model(**inputs, output_hidden_states=True) layer3_pt = outputs.hidden_states[3].cpu().numpy() # [1, 10, 4096] # TF端:用layer3_pt作为输入,验证下一层输出一致性 layer3_tf = tf.constant(layer3_pt, dtype=tf.float32) layer4_tf = tf_model.layers[4](layer3_tf) # 第4层Transformer # 比对:layer4_tf vs outputs.hidden_states[4].cpu().numpy()4. 权重转换:从.bin到.ckpt的七步炼金术与三个维度陷阱
4.1 PyTorch权重结构解析:为什么pytorch_model.bin不能直接tf.train.load_checkpoint
DeepSeek-7B的pytorch_model.bin是state_dict字典,键名为model.layers.0.self_attn.q_proj.weight,而TF SavedModel要求变量名encoder/layer_0/attention/q_proj/kernel。文档第28页指出:92%的转换失败源于键名映射错误,而非维度问题。核心映射规则:
| PyTorch键名模式 | TF变量名模式 | 转换逻辑 |
|---|---|---|
model.embed_tokens.weight | embedding/token_embedding/kernel | 前缀替换+后缀标准化 |
model.layers.0.self_attn.q_proj.weight | decoder/layer_0/attention/q_proj/kernel | 层数提取+模块重命名 |
model.norm.weight | decoder/final_layernorm/gamma | 归一化层特殊处理 |
4.2 维度对齐:[out_features, in_features]到[in_features, out_features]的致命翻转
PyTorch线性层权重形状为[out_features, in_features],TF默认为[in_features, out_features]。但DeepSeek的q_proj权重在PyTorch中是[5120, 4096](7B模型),TF需转为[4096, 5120]。文档第41页强调:必须用np.transpose(weight, (1,0))而非weight.T,因为weight.T在torch.float16下会触发隐式类型转换错误:
# 错误:会丢失精度 tf_weight = weight.T.numpy() # float16 -> float32隐式转换 # 正确:保持dtype tf_weight = np.transpose(weight.numpy(), (1, 0)).astype(weight.dtype)4.3 数据类型转换:bfloat16在TF中的“幽灵精度”问题
PyTorch的bfloat16在TF 2.12中无原生支持,tf.bfloat16仅在TPU上可用。GPU上必须转为tf.float32,但文档第42页发现:直接weight.to(torch.float32)会引入1e-3级误差。解决方案是使用torch.float32中间态:
# PyTorch端:用float32作为转换桥梁 pt_weight_bf16 = state_dict["model.layers.0.mlp.gate_proj.weight"] # bfloat16 pt_weight_fp32 = pt_weight_bf16.to(torch.float32) # 精确转换 # TF端:转为float32变量 tf_weight = tf.Variable( initial_value=pt_weight_fp32.numpy(), dtype=tf.float32, name="decoder/layer_0/mlp/gate_proj/kernel" )4.4 权重分片合并:处理pytorch_model-00001-of-00002.bin的原子操作
DeepSeek-67B权重被分片为多个.bin文件。文档第31页警告:不能简单torch.load后dict.update(),会丢失_metadata导致load_state_dict(strict=True)失败。正确方法是用shard_utils:
from transformers.modeling_utils import shard_checkpoint # 加载所有分片 shards = [] for shard_file in ["pytorch_model-00001-of-00002.bin", "pytorch_model-00002-of-00002.bin"]: shards.append(torch.load(shard_file, map_location="cpu")) # 合并分片(保留_metadata) merged_state_dict = {} for shard in shards: merged_state_dict.update(shard) # 修复_metadata(关键!) merged_state_dict["_metadata"] = {"total_size": sum(v.numel() for v in merged_state_dict.values())} # 保存为单个bin供TF读取 torch.save(merged_state_dict, "deepseek-67b-merged.bin")4.5 自定义转换工具:deepseek2tf命令行工具的核心逻辑
文档第38页开源的deepseek2tf工具,其核心是WeightConverter类:
class WeightConverter: def __init__(self, pt_path: str, tf_path: str): self.pt_state_dict = torch.load(pt_path, map_location="cpu") self.tf_vars = {} # {tf_var_name: np.ndarray} def convert(self): for pt_key, pt_weight in self.pt_state_dict.items(): if pt_key == "_metadata": continue tf_key = self._map_key(pt_key) # 键名映射 tf_weight = self._convert_weight(pt_weight) # 维度+类型转换 self.tf_vars[tf_key] = tf_weight def save_as_ckpt(self): # 创建TF Checkpoint checkpoint = tf.train.Checkpoint(**{ name: tf.Variable(value, name=name) for name, value in self.tf_vars.items() }) checkpoint.write(f"{self.tf_path}/model.ckpt")4.6 参数校验:torch.allclose在TF中的等效实现
文档第44页提供TF端一致性验证函数:
def tf_allclose(a: tf.Tensor, b: tf.Tensor, atol=1e-5) -> bool: """TF版allclose,避免tf.math.reduce_all返回scalar""" diff = tf.abs(a - b) max_diff = tf.reduce_max(diff) return bool(max_diff.numpy() <= atol) # 使用示例 pt_output = ... # PyTorch推理输出 tf_output = ... # TF推理输出 assert tf_allclose(pt_output, tf_output, atol=1e-4), "权重转换偏差超限"5. 避坑:跨框架迁移中五个血泪教训与即时修复方案
5.1 现象:TensorFlow推理输出全为nan
原因:PyTorch权重中存在inf值(常见于训练中断的checkpoint),TF加载后触发nan传播。
解决:在转换前清洗PyTorch权重:
# 清洗inf/nan for k, v in pt_state_dict.items(): if torch.is_floating_point(v): v = torch.where(torch.isfinite(v), v, torch.zeros_like(v)) pt_state_dict[k] = v5.2 现象:tf.function装饰后训练速度下降50%
原因:DeepSeek的动态padding逻辑(如tf.pad)在@tf.function中触发trace重编译。
解决:改用tf.data.Dataset.padded_batch预处理:
# 错误:在model.call()中动态pad def call(self, inputs): padded = tf.pad(inputs, [[0,0],[0,max_len-tf.shape(inputs)[1]]]) # 触发re-trace # 正确:在Dataset中pad dataset = dataset.padded_batch( batch_size=8, padded_shapes=([None, 4096], [None]), # 预设最大长度 padding_values=(0, 0) )5.3 现象:PyTorch与TF的LayerNorm输出相差1e-2
原因:PyTorchnn.LayerNorm默认elementwise_affine=True,TFtf.keras.layers.LayerNormalization默认center=True, scale=True,但epsilon值不同(PyTorch=1e-5,TF=1e-3)。
解决:TF端显式设置epsilon=1e-5:
# TF LayerNorm必须指定epsilon ln = tf.keras.layers.LayerNormalization( axis=-1, epsilon=1e-5, # 与PyTorch对齐 center=True, scale=True )5.4 现象:torch.compile加速后TF转换失败
原因:torch.compile会修改state_dict键名(如添加_orig_mod.前缀),导致键名映射失效。
解决:转换前禁用compile或提取原始模块:
# 方案1:转换前解除compile if hasattr(model, '_compiled_module'): model = model._orig_mod # 方案2:直接访问原始模块 original_model = model._orig_mod if hasattr(model, '_orig_mod') else model state_dict = original_model.state_dict()5.5 现象:分布式训练中TF的MultiWorkerMirroredStrategyloss为0
原因:DeepSeek的CrossEntropyLoss在TF中需用tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),而PyTorch用nn.CrossEntropyLoss,二者数值计算路径不同。
解决:TF端手动实现PyTorch风格loss:
def pytorch_ce_loss(y_true, y_pred): # y_pred: [batch, seq, vocab], y_true: [batch, seq] y_true = tf.cast(y_true, tf.int32) # 手动计算log_softmax + nll_loss log_probs = tf.nn.log_softmax(y_pred, axis=-1) nll_loss = tf.nn.sparse_softmax_cross_entropy_with_logits( labels=y_true, logits=y_pred ) return tf.reduce_mean(nll_loss)6. 推理一致性验证:用torch.allclose构建TF的黄金标尺
6.1 构建跨框架验证流水线:从单层到端到端的四阶测试
文档第167页定义的验证金字塔:
| 测试层级 | 输入 | 输出 | 通过标准 | 工具 |
|---|---|---|---|---|
| Layer级 | 随机张量 | 单层输出 | tf_allclose(pt_out, tf_out, atol=1e-5) | tf.test.TestCase |
| Module级 | tokenized input | hidden_states | 各层hidden_states逐层比对 | 自定义DeepSeekModuleTester |
| Model级 | prompt string | logits | argmax(logits)一致率≥99.9% | transformers.pipeline |
| End2End级 | 用户query | 生成文本 | BLEU-4 ≥ 0.998(vs PyTorch baseline) | sacrebleu |
6.2 Layer级验证:以DeepSeekAttention为例的原子测试
class TestDeepSeekAttention(tf.test.TestCase): def test_attention_output_consistency(self): # 生成PyTorch参考输出 pt_attn = DeepSeekAttention(config) # PyTorch版 pt_q = torch.randn(2, 128, 4096) pt_k = torch.randn(2, 128, 4096) pt_v = torch.randn(2, 128, 4096) pt_out, _ = pt_attn(pt_q, pt_k, pt_v) # 生成TF输出 tf_attn = TFDeepSeekAttention(config) # TF版 tf_q = tf.constant(pt_q.numpy()) tf_k = tf.constant(pt_k.numpy()) tf_v = tf.constant(pt_v.numpy()) tf_out = tf_attn(tf_q, tf_k, tf_v) # 严格比对 self.assertAllClose( pt_out.detach().numpy(), tf_out.numpy(), atol=1e-5, msg="Attention层输出不一致" ) if __name__ == "__main__": tf.test.main() # 运行TF测试套件6.3 Model级验证:用Hugging Face pipeline统一接口
# PyTorch pipeline pt_pipeline = pipeline( "text-generation", model="deepseek-ai/deepseek-llm-7b-base", tokenizer="deepseek-ai/deepseek-llm-7b-base", torch_dtype=torch.bfloat16, device_map="auto" ) # TF pipeline(需先将模型转为SavedModel) tf_pipeline = TFPipeline( model_dir="/path/to/tf_savedmodel", # 文档第177页导出的SavedModel tokenizer="deepseek-ai/deepseek-llm-7b-base" ) # 批量测试100个prompt prompts = ["Explain quantum computing", "Write Python code for quicksort"] for prompt in prompts: pt_result = pt_pipeline(prompt, max_new_tokens=20)[0]["generated_text"] tf_result = tf_pipeline(prompt, max_new_tokens=20)[0]["generated_text"] # 计算BLEU-4 bleu_score = sentence_bleu([pt_result.split()], tf_result.split()) self.assertGreater(bleu_score, 0.998)6.4 End2End验证:生产环境下的实时监控方案
文档第119页部署的ConsistencyMonitor服务:
# 实时监控脚本 class ConsistencyMonitor: def __init__(self, pt_model, tf_model): self.pt_model = pt_model self.tf_model = tf_model self.metrics = { "output_diff_mean": [], "output_diff_std": [], "argmax_match_rate": [] } def monitor_step(self, input_ids): # 同时运行PT和TF pt_logits = self.pt_model(input_ids).logits tf_logits = self.tf_model(input_ids).logits # 计算指标 diff = np.abs(pt_logits.numpy() - tf_logits.numpy()) self.metrics["output_diff_mean"].append(np.mean(diff)) self.metrics["output_diff_std"].append(np.std(diff)) self.metrics["argmax_match_rate"].append( np.mean(np.argmax(pt_logits.numpy(), axis=-1) == np.argmax(tf_logits.numpy(), axis=-1)) ) # 告警阈值 if np.mean(diff) > 1e-3: send_alert("Output diff too high!") def report(self): return { "mean_diff": np.mean(self.metrics["output_diff_mean"]), "std_diff": np.mean(self.metrics["output_diff_std"]), "match_rate": np.mean(self.metrics["argmax_match_rate"]) } # 在生产API中嵌入 @app.route("/generate", methods=["POST"]) def generate(): data = request.json input_ids = tokenizer(data["prompt"], return_tensors="pt").input_ids monitor.monitor_step(input_ids) return {"text": tf_model.generate(input_ids)}从那以后我每次交付跨框架模型,都强制走一遍这四阶验证——哪怕客户说“只要能跑就行”。因为线上一个nan输出,可能让整个金融风控模型误判百万订单。希望帮到你。
本文还有配套的精品资源,点击获取