PyTorch量化线性函数报形状无效错误原因排查求助
问题原因及解决方案
核心问题
PyTorch的QF.linear(量化线性函数)与F.linear的输入维度要求完全不同:
F.linear支持任意≥2D的输入(自动将额外维度视为batch的一部分)QF.linear仅接受2D量化张量(形状为[batch_size, in_features]),同时要求权重为2D量化张量,且当前版本对CUDA量化张量的支持有限。
你的输入是3D张量(2,512,4096),直接传入QF.linear会导致内部维度解析错误,触发形状不匹配的RuntimeError。
修复步骤及代码修改
1. 调整输入维度,适配QF.linear的要求
将3D输入flatten为2D,计算完成后再恢复原维度结构。
2. 切换到CPU执行量化操作
当前PyTorch 2.4.1中,QF.linear对CUDA量化张量的支持不完善,建议先将张量转回CPU再量化。
3. 确保权重量化的正确性
保持权重形状为[out_features, in_features],量化类型符合QF.linear的要求(输入quint8,权重qint8)。
修改后的代码如下:
import torch import torch.ao.nn.quantized.functional as QF import torch.nn.functional as F loaded_data = torch.load('ffn_w3_example.pt') w3 = loaded_data['ffn_w3'].type(torch.float16) x = loaded_data['x'].type(torch.float16) original_out = loaded_data['out'].type(torch.float16) out_noquant = F.linear(x, w3) def scale_zpt_compute(inp, Q_MAX, Q_MIN): scale = (inp.max() - inp.min()) / (Q_MAX - Q_MIN) zero_point = torch.round(- torch.min(inp) / scale) + Q_MIN return scale, zero_point # 量化权重为qint8,保持形状(14336, 4096),转回CPU处理 Q_MAX = 127.0 Q_MIN = -127.0 scale_w, zero_point_w = scale_zpt_compute(w3.float(), Q_MAX, Q_MIN) w3_q = torch.quantize_per_tensor(w3.float().cpu(), scale_w, zero_point_w, dtype=torch.qint8) # 量化输入为quint8,先flatten为2D,再量化 Q_MAX = 255.0 Q_MIN = 0.0 batch_size, seq_len, in_features = x.shape x_flat = x.float().cpu().reshape(-1, in_features) # 形状变为(2*512, 4096) scale_x, zero_point_x = scale_zpt_compute(x_flat, Q_MAX, Q_MIN) x_q = torch.quantize_per_tensor(x_flat, scale_x, zero_point_x, dtype=torch.quint8) # 执行量化线性计算,再恢复原维度 out_q_flat = QF.linear(x_q, w3_q) out_q = out_q_flat.dequantize().reshape(batch_size, seq_len, -1).to(torch.float16) # 验证结果(可选) print(f"原输出形状: {out_noquant.shape}") print(f"量化输出形状: {out_q.shape}")
额外说明
- 如果需要在CUDA上执行量化线性操作,建议使用PyTorch的动态量化或**量化感知训练(QAT)**流程,直接使用
quantize_per_tensor手动量化的方式在CUDA上兼容性较差。 - 动态量化示例:可以使用
torch.ao.nn.quantized.dynamic.Linear来替代手动调用QF.linear,它会自动处理维度和设备问题。
内容的提问来源于stack exchange,提问作者hafezmg48
相关产品推荐
相关产品推荐

