ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch模型权重量化与FPGA部署:从浮点到定点补码的实践

PyTorch模型权重量化与FPGA部署:从浮点到定点补码的实践 前阵子接手了个活儿把 PyTorch 训练好的一个几层卷积网络搬到 FPGA 上做实时小目标检测。模型不大参数量几十万但问题很现实——训练完的权值矩阵全是 float32而 FPGA 里我最想用的 BRAM 和 DSP 硬核都不吃这一套。浮点乘加在 FPGA 上要么占用大量逻辑资源去做软核要么只能用有限的 DSP 硬核硬扛存储更是直接暴露在片上 RAM 容量面前。于是就有了这次“将 PyTorch 的权值矩阵量化为定点数补码并导入 FPGA”的完整折腾记录。这篇内容适合正在做边缘推理加速、想把训练好的模型落到 FPGA 上或者对“PyTorch 模型权重怎么变成硬件能读的数据”这个环节有疑惑的同学。我会把浮点转定点的原理、为什么非得用补码、PyTorch 端怎么批量提取和量化权值、怎么生成 .coe/.mem 文件、FPGA 端怎么取数和闭环验证这些步骤按我实操的顺序一条线讲完。文中所有代码都是我可直接运行的版本你拿到后改一改路径和层名就能套用。1. 为什么非要把权值矩阵变成定点补码1.1 浮点数的存储代价远比你想象的高PyTorch 默认用 float32 保存模型权值一个权重占用 32 bit。FPGA 里的片上存储资源是 BRAM容量从几十 KB 到几 MB 不等。一个小型卷积网络动辄几十万参数全用 float32 存的话光权值就要占掉几 MB很多中低端芯片根本放不下更别提还要给激活值、中间结果留空间。更麻烦的是计算资源。FPGA 里的 DSP 硬核通常使用整数乘法器处理 8bit×8bit、16bit×16bit 很高效但直接做 float32 乘法需要额外搭浮点运算单元逻辑资源消耗极大、时序也很容易跑不上去。所以实际做推理加速时最常规的思路就是把浮点权值压缩成定点数用更少的 bit 表示近似相同的数值然后让 DSP 直接处理整数乘加。这样做的收益很直观float32 存一个权重要 4 字节量化成 int8 后只需要 1 字节存储直接省 4 倍计算时 DSP 硬核跑 8bit 乘法速度和资源开销都比浮点软核友好太多。用一个 3×3 卷积、输入 32 通道、输出 64 通道的层算一下权值数量是 3×3×32×6418432 个float32 存储需要 72KBint8 只需要 18KB这个差距在 BRAM 里非常可观。1.2 补码是 FPGA 里数字电路最喜欢的“格式”聊定点数之前必须先说清楚补码。数字电路里没有减法器只有加法器所以负数必须用一种“减法能变成加法”的编码方式这就是补码存在的意义。补码的定义很简单对于 n bit 有符号数最高位是符号位负数的补码等于其绝对值的原码按位取反再加 1。举个例子用 8bit 表示 -5。5 的原码是 0000_0101取反得 1111_1010再加 1 得 1111_1011这就是 -5 的补码。好处在于硬件只需要一个加法器就能完成加减法比如 3 (-5) 等价于 0000_0011 1111_1011 1111_1110这个结果正好是 -2 的补码。这个特性让补码成了所有处理器和 DSP 硬核默认采用的有符号数格式。FPGA 里做模型推理时权值有正有负如果用原码或者反码电路里每个乘法都要额外处理符号位很不划算。用补码的话直接把量化后的整数塞给 DSP 就行硬件端不需要关注“这个数到底是不是负数”符号信息天然被编码在位模式里。这也是为什么标题里特别强调“补码”——它不是可选项而是硬件端能高效工作的前提。1.3 量化误差可控关键是定标把 float32 变成 int8 一定会有精度损失但工程上完全可控。一个训练好的网络权值分布通常很集中大部分权重绝对值都在 0.01 到 0.5 之间只有少部分较大。我们要做的就是从这些权值中找到合适的“定标参数”让量化后的整数能尽可能精确地表示原始浮点值。定标参数主要有两个一个是缩放因子 scale表示“整数走一步对应浮点数走多少”另一个是零点 zero_point表示“整数 0 对应浮点数多少”。如果权值分布大致关于 0 对称我们通常用对称量化zero_point 直接就是 0这样实现最简单硬件端完全不需要做零点偏移的补偿。量化误差的大小取决于两点位宽和定标是否合适。位宽越大误差越小但存储和计算成本越高定标偏差过大会导致很多权重被截断到同一个整数误差迅速放大。后面我会给出具体的量化公式和代码这部分会看得更明白。2. 量化方案设计2.1 对称量化和非对称量化的区别量化方案有对称量化和非对称量化两种。对称量化假设浮点数值域关于 0 对称映射关系是q round(r / scale) scale max_abs / (2^(bits-1) - 1)这里的 max_abs 是整个矩阵里绝对值最大的那个数。比如 8bit 量化正数最大值是 127负数最小值是 -128。因为是镜像对称所以 scale max_abs / 127。量化后最小负整数是 -128对应浮点数 -128 × scale会比 -max_abs 稍大一点。非对称量化则是单独记录浮点最小值 min_val 和最大值 max_val用整数区间完整映射浮点区间公式变成scale (max_val - min_val) / (2^bits - 1) zero_point round(-min_val / scale) q round(r / scale) zero_point非对称量化精度更高尤其是当数据分布偏到一边时比如激活函数 ReLU 的输出全是非负的用非对称量化能充分利用整数的全部取值区间。但代价是硬件端要多做一次零点减法计算复杂度增加。对于大多数卷积层的权值因为初始化、正则化等原因分布通常比较对称用对称量化就够而且实现更简洁、错误更少。我建议第一阶段先无脑用对称量化如果发现某些层精度损失特别严重再单独给这些层切到非对称量化。2.2 位宽和 Q 格式怎么选选位宽本质上是平衡存储、计算效率和精度三者的关系。实际项目里最常用的是 8bit原因很现实FPGA 的 DSP 硬核通常一次能处理 8bit×8bit 乘法BRAM 按 8bit/16bit/32bit 组织效率也最高。有些对精度敏感的网络会用到 16bit比如第一层输入或者最后的全连接层但整体上 8bit 是性价比首选。选定 bit 宽度之后还要确定小数点在哪个位置这就是 Q 格式。Qm.n 表示用 m bit 表示整数部分含符号位n bit 表示小数部分总位宽 mn。比如 Q8.8 表示 16bit 里有 8bit 整数、8bit 小数取值范围是 -32768/256 到 32767/256也就是 -128 到 127.99609375精度为 1/256。选择 Q 格式的关键是看权值的动态范围。如果一层权值最大绝对值是 0.25用 8bit 表示那么最理想的情况是把这个范围映射到整个整数区间即 scale 0.25/127 ≈ 0.00197。对纯整数表示来说这等价于把小数点在 8bit 整数里的位置调整到“当前数值范围内最精细”的位置也就是让定点数的最低位满足精度需求。我实际使用时的经验是不一定要把 Q 格式固定成某个全局参数可以按层统计权值范围并记录各自的 scale导出时给每层带上自己的 scale 值。这样比强制全网络共用一个 Q 格式更省 bit精度也好一些。FPGA 端只需要从 BRAM 的固定地址读出每层的 scale然后统一做一次右移对齐。2.3 垃圾进垃圾出量化前先检查权值分布在写任何量化代码之前我强烈建议先跑一段脚本看一下各层权值的统计信息。这一步不是形式主义而是防止量化后精度崩了再回头排查时找不到方向。具体做法也很简单加载完 state_dict 后对每个 weight 张量打印 shape、min、max、mean、std、absmax。很多时候你会发现全连接层最后一层的权值范围比其他卷积层大好几倍如果按全局统一 scale 量化小数值的卷积层会被严重截断精度损失巨大。更合理的做法是逐层定标、逐层导出。3. PyTorch 端权值提取与量化实操3.1 从 state_dict 里把权值捞出来先加载模型权重。PyTorch 里所有可训练参数都挂在 model.state_dict() 下key 是层名value 是 Tensor。加载完模型后我建议按层类型过滤因为我们需要的是 Conv2d 和 Linear 的 weightBatchNorm 层的 mean、variance 虽然也要用到推理里但它们有另外的处理方式这里先不混在一起。提取权值有几个坑。第一state_dict 的 key 是字符串网络里不同层会带 numbered 前缀最好自己打印一遍确认 key 的命名规则。第二卷积层的权值 shape 是 (out_channels, in_channels, kh, kw)导出成二进制数据时要按存储顺序逐元素展平C 语言风格的行优先顺序要和 Verilog 里读地址的顺序一致。第三bias 也是可训练参数后面做推理对齐时必须一起量化否则卷积计算结果会对不上。我通常会把代码写成支持传入 [layer_prefix] 列表只量化指定层这样便于逐层调试import torch # 假设 model 是已经定义好的网络结构且已经加载了权重 state_dict model.state_dict() # 看一下所有 key for k in state_dict.keys(): print(k, tuple(state_dict[k].shape))实际项目里我会用正则匹配把 conv 和 linear 的 weight 区分出来这样批量处理非常方便不会漏层也不会把 BN 参数误当成卷积权值导出。3.2 从浮点权值到定点补码的完整转换代码下面这套代码是我在多个项目里反复使用的“标准化导出器”。它做的事情很清晰遍历指定层读取浮点权值计算 scale量化成整数转成补码位模式最后统一输出。import torch import numpy as np def quantize_symmetric(weight_fp, bits8): 对称量化把浮点权值转成有符号整数 返回量化后的整数张量 q_weight 和 scale qmax 2 ** (bits - 1) - 1 # 8bit - 127 qmin -(2 ** (bits - 1)) # -128 max_abs weight_fp.abs().max().item() if max_abs 0: max_abs 1e-12 scale max_abs / qmax q_weight torch.round(weight_fp / scale) q_weight q_weight.clamp(qmin, qmax) # 必须截断防止溢出 return q_weight.to(torch.int64), scale def to_twos_complement_uint(q_weight, bits8): 关键步骤把有符号整数转成补码对应的无符号整数位模式 原理补码在截断到 bits 位时等于原数对 2^bits 取模 mask (1 bits) - 1 # 8bit - 0xFF return q_weight mask def export_layer_weights(layer_name, weight_fp, bits8, output_pathNone): q_weight, scale quantize_symmetric(weight_fp, bits) uint_weight to_twos_complement_uint(q_weight, bits) # 展平成一维数组行优先顺序 flat uint_weight.cpu().numpy().flatten().astype(np.uint64) # 记录关键元信息 meta { layer: layer_name, shape: list(weight_fp.shape), bits: bits, scale: scale, count: flat.shape[0], } print(f[导出] {layer_name}: shape{meta[shape]}, fbits{bits}, scale{scale:.6f}, count{flat.shape[0]}) if output_path: np.save(output_path, flat) # 同时把元信息写到文本方便 FPGA 端读取 with open(output_path .meta.txt, w) as f: for k, v in meta.items(): f.write(f{k}: {v}\n) return flat, meta # 用法示例 flat, meta export_layer_weights( layer_nameconv1.weight, weight_fpstate_dict[conv1.weight], bits8, output_pathweights_conv1.npy )这段代码里最重要的就是to_twos_complement_uint这个函数。它的原理是一个负数在 n bit 补码表示下等于这个数加上 2^n 后对应的无符号整数。所以“转补码”这种听起来工程味儿十足的事情在 Python 里其实就是一次与掩码的位运算。比如 -128 转 8bit 补码等价于 (-128) 0xFF 0x80也就是十六进制 0x80完全正确。3.3 生成 .coe 文件和 .mem 文件PyTorch 这边把补码整数算出来后下一步是把它变成 FPGA 工具链能识别的初始化文件。最常用的是 Xilinx Vivado 的 COE 文件和通用的 MEM 文件。COE 文件主要用于 Block Memory Generator IP 核初始化MEM 文件则配合 Verilog 的$readmemh使用。COE 文件的格式很简单先声明进制再用逗号分隔数据最后以分号结束。我写过一段生成 COE 的代码逻辑如下def write_coe(filepath, data_flat, bits8): data_flat 是无符号整数数组已经转成补码位模式 输出 Vivado 可用的 COE 文件 hex_per_line 16 # 根据位宽决定十六进制位宽8bit 用 2 个 hex16bit 用 4 个 hex hex_width (bits 3) // 4 with open(filepath, w) as f: f.write(memory_initialization_radix16;\n) f.write(memory_initialization_vector\n) total len(data_flat) for i in range(0, total, hex_per_line): chunk data_flat[i:ihex_per_line] line , .join(f{v:0{hex_width}x} for v in chunk) if i hex_per_line total: f.write(line ;\n) else: f.write(line ,\n)MEM 文件更简单每行一个十六进制数常用于$readmemh。生成方式如下def write_mem(filepath, data_flat, bits8): hex_width (bits 3) // 4 with open(filepath, w) as f: for v in data_flat: f.write(f{v:0{hex_width}x}\n)使用经验如果你还在仿真阶段用 MEM 文件最方便直接在 testbench 里读。如果已经准备综合上板用 COE 文件初始化 BRAM IP 核更省心。两者都不复杂但在导出时一定要记录好每个文件对应的层名、位宽、数据总量和地址偏移后面在 FPGA 端查找问题会救命。3.4 量化误差验证导出前先算一笔账不要急着把文件扔给 FPGA导出前先在 PyTorch 里做一次“定点回环仿真”把量化后的整数反量化回浮点看看和原始浮点权值差多少。操作非常简单def deprecated_quant_error(weight_fp, bits8): q_weight, scale quantize_symmetric(weight_fp, bits) dequant_weight q_weight.double() * scale abs_err (dequant_weight - weight_fp).abs() rel_err abs_err / (weight_fp.abs() 1e-12) mae abs_err.mean().item() mape rel_err.mean().item() cos_sim torch.nn.functional.cosine_similarity( dequant_weight.flatten().double().unsqueeze(0), weight_fp.flatten().double().unsqueeze(0) ).item() return mae, mape, cos_sim我踩过一次最深刻的坑就是有一个全连接层的权重绝对值非常小最大值只有 0.002而网络里其他层最大值接近 0.5。如果整个网络共用一个 scale那小数值层量化后几乎全部变成 0余弦相似度直接掉到 0.93 以下。后来改成逐层定标每层用自己的 max_abs 算 scale余弦相似度立刻恢复到 0.998 以上。这个检查步骤非常便宜但能省掉一整天的硬件联调时间。4. FPGA 端导入与计算对齐4.1 Vivado 里用 Block Memory Generator 加载 COE在实际 FPGA 工程里我会把每层权值存到一个独立 BRAM 或 ROM 里。Xilinx 平台最常用的办法是用 Block Memory Generator IP 核配置成 Single Port ROM位宽填 bits通常 8 或 16深度填权值个数然后在 Other Options 里加载 COE 文件。这一步记得要用“Load COE file”Vivado 会把 COE 里的十六进制数据直接烧进 BRAM 初始化内容里。如果位宽和 COE 文件数据宽度不一致工具通常并不会报错但读取时数据会错位得莫名其妙。我建议生成 COE 时就严格按最终 BRAM 位宽输出别在 Vivado 里做二次转换。一个特别容易踩的坑是 COE 文件最后一行的分号。Vivado 对格式非常敏感数据行最后必须以英文分号结尾少一个都会导致加载失败。我在写write_coe函数时特地把分号逻辑写严了如果你手写文件一定检查最后一行。4.2 从 BRAM 读出后如何正确解释成有符号数BRAM 在 Verilog 里默认是无符号的存储载体。如果权值 8bit 补码是 0x80也就是 -128直接把它当 unsigned 参与计算就错了会变成 128。解决办法有两个要么在 RAM 声明时用signed要么在读取时用$signed()做符号扩展。我常用的方式是第一种在模块里直接这样写reg [7:0] weight_ram [0:DEPTH-1]; wire signed [7:0] weight_s $signed(weight_ram[addr]);这样在乘加运算里weight_s就能正确带着符号位参与 DSP 计算。还有一种方式是直接把 RAM 声明成reg signed [7:0] weight_ram [0:DEPTH-1];但某些 BRAM IP 核生成的原语不支持 signed 数组或者综合工具会自动把它转成 unsigned。为了保证代码通用性我更推荐读取时用$signed转换。还有一个小细节8bit 补码的范围是 -128 到 127非对称。如果某一层量化后所有数值正好是 -128那说明 scale 算的时候溢出风险很高建议检查一下是否发生了饱和截断。反之如果数据集里几乎没有负的极值那说明定标还可以更激进一些。4.3 矩阵乘/卷积计算时的对齐细节量化后的整型权值和整型输入相乘结果实际上还是“定点数”只不过位宽扩大了。比如两个 8bit 数相乘得到 16bit 结果这个结果的定标不是两个 scale 简单相乘就完事还要考虑硬件里小数点位置的隐含对齐。假设输入 feature map 用 8bit 定点scale 为 s_i权值用 8bit 定点scale 为 s_w那么乘法结果的数学含义是浮点值(int_input * s_i) * (int_weight * s_w)等于整数乘积乘以s_i * s_w。但硬件里只能算整数乘积需要我们在累积完所有乘法后对累加结果做一次移位或者乘法把小数部分重新归一化到输出需要的格式。具体做法有两种。一种是量化输出时直接把累加结果除以s_i * s_w再量化到输出位宽另一种是提前把所有定标参数设成 2 的幂次这样除法就是一次右移FPGA 端成本极低。第二种方法在工程里更常用比如把 scale 强制设计成 2 的负 k 次方用移位就能完成反量化。但要注意如果量化范围很小强制用 2 的幂次方会让部分动态范围损失需要做取舍。我在 FPGA 端做卷积时会把 im2col 展开后直接用乘累加阵列完成计算最后统一做一次右移。这里的核心原则是每一层都必须记录自己的 scale并提前算好输出应该右移多少 bit。如果这个“右移量”不匹配网络推理结果会出现系统性偏差而且越往后累积越严重。4.4 闭环验证怎么确认 FPGA 算的和 PyTorch 定点仿真一致把数据导入 FPGA 后第一件事不是直接跑完整网络而是做小规模的“单层闭环验证”。我会在 PyTorch 里把第一层的输入也量化为定点数然后手动做一次定点卷积拿到参考输出FPGA 端只加载第一层权值用同一份输入数据跑一次对比结果。在 FPGA 仿真里我会用 testbench 加载输入和权值跑完一次卷积后把结果写回文本文件然后在 Python 里对比。对比标准很简单只要 FPGA 输出和 PyTorch 定点仿真的整数结果完全一致就说明数据导入和算术链路都对了。然后再对比 PyTorch 浮点输出计算精度损失有多大这一步才衡量量化方案到底行不行。如果第一层对不上重点排查方向是权值地址错位、字节序反了、符号扩展没做、scale 右移量不对。如果第一层对上了但后面层对不上问题大概率出在层间数据传递格式没有对齐比如累计结果没有正确截断到下一层输入位宽。5. 实操中遇到的常见问题与排查技巧5.1 常见问题速查表现象可能原因解决办法COE 加载后 Vivado 报错文件结尾分号缺失、进制声明错误检查最后一行是;确认 radix 是 16BRAM 读出的数据全是 0COE 没有被加载进 IP、地址越界重新生成 IP核对深度确认初始化选项勾选负数权值全变成很大的正数没有做$signed()符号扩展读取 BRAM 数据后加$signed()某一层输出整体偏大或偏小右移量不对scale 没对齐回到 PyTorch 复核该层 scale算出正确右移位数量化后精度损失爆表层间动态范围差异大共用了同一个 scale改为逐层定标每层独立 scale网络前边正确、后边全乱层间数据格式没截断位宽膨胀每一层输出强制截断到下一层输入位宽内存占用爆炸位宽选太大或 BRAM 实例化过多优先 8bit必要时只对敏感层用 16bit综合后时序不收敛乘累加阵列过长对累加路径做流水线切分5.2 字节序和地址对齐的典型坑早期我吃过一次大亏导出权值后直接在 FPGA 端用连续地址读取结果每隔几个数就错一个。后来发现是 Python 里用numpy保存成.npy文件然后转成二进制时默认使用了小端序而 BRAM 初始化文件是按逐字节顺序写的我没有统一字节序规则。现在我的习惯是无论最后导出 COE 还是 MEM都在导出函数里明确指定数据的“逻辑位宽”和“存储字节宽”。如果权值是 8bit一个权值就是一个字节不存在字节序问题如果权值是 16bit就一定要约定好高位在前还是低位在前并在注释里写清楚。我项目里统一用大端模式高字节放低地址这样和 Vivado 默认显示顺序一致排查起来比较舒服。5.3 如何快速定位是量化问题还是硬件问题整套流程如果出错先别急着改 FPGA 代码。推荐按时间成本从低到高排查第一在 PyTorch 里做定点仿真用同一份量化代码模拟硬件上的整数乘加和右移截断。这一步能筛掉大部分算法定标问题比如 scale 算错、截断溢出、右移量不对。第二做 RTL 仿真并和定点仿真逐时钟对比如果 RTL 仿真结果和 PyTorch 定点模拟的结果不一致问题在硬件逻辑。第三再上板测试如果板级结果和 RTL 仿真一致说明硬件链路都通了剩下只是输入输出接口的时序问题或者固定偏移。这个排查顺序我屡试不爽。很多朋友一上来就抓逻辑分析仪去调板级信号绕了一大圈才发现 PyTorch 导出的 scale 就是个错的。先在软件里把“定点闭环”做到 100% 匹配再上硬件链路会顺很多。6. 一些总结性经验这套流程我已经在不同项目里跑过好多遍从最初手工写脚本一步步处理到后来封装成“一键导出”工具核心思路基本没变浮点权值矩阵 → 逐层定标 → 量化成整数 → 转成补码位模式 → 导出为 COE/MEM 文件 → FPGA 端按地址读入 → 按 scale 对齐反量化。每一步都不算难但每一步都有细节坑。我个人在实际项目中最大的体会是定标参数的记录和传递是整套系统的“命门”。只要导出的 .meta.txt 里记录了每层的 scale、位宽、数据个数FPGA 端的所有对齐问题都能顺着这个清单快速定位。相反如果只导出裸数据丢了 scale后面任何验证环节都会变成无头苍蝇。最后再分享一个小技巧别把 8bit 当成唯一解。实际测试时我经常遇到某个网络所有层都用 8bit 跑得好好的但最后一层全连接输出掉点明显。这时候只需要把输出层单独设成 16bit、其他层保持不变精度立刻就能回来资源开销增加却很小。逐层定制位宽和 scale远比“一刀切”的全局量化效果好。如果你也在做 PyTorch 模型到 FPGA 的部署建议先把整个流程跑通哪怕先用一个很小的网络。跑通了之后再逐步加层、加功能、调精度。这条路我第一次走的时候花了差不多一周现在有了这套方法和代码基本上一两天就能出结果。希望这篇记录能帮你少走点弯路。
返回列表