归一化输入输出张量时出现NaN结果的技术求助
问题分析与解决方案
核心原因
你的张量conts_total中存在NaN值,导致调用.min()和.max()时返回NaN,最终归一化计算出现全NaN结果。
步骤1:确认并定位NaN值
先验证张量中是否存在NaN:
# 检查是否有NaN print(torch.isnan(conts_total).any()) # 统计NaN总数量 print(torch.isnan(conts_total).sum().item()) # 查看哪些样本包含NaN nan_rows = torch.where(torch.isnan(conts_total).any(dim=1))[0] print(nan_rows)
步骤2:处理NaN值
根据数据场景选择以下方式之一:
- 移除含NaN的样本(样本量充足时优先):
# 保留所有不含NaN的行 conts_total = conts_total[~torch.isnan(conts_total).any(dim=1)] - 用列均值填充(样本量有限时):
for col_idx in range(conts_total.shape[1]): col = conts_total[:, col_idx] # 获取该列非NaN值的均值 col_mean = col[~torch.isnan(col)].mean() # 填充NaN col[torch.isnan(col)] = col_mean - 用固定值填充(如0,需结合特征含义判断):
conts_total = torch.nan_to_num(conts_total, nan=0.0)
步骤3:正确实现Min-Max归一化
机器学习中通常按特征(列)单独归一化,而非全局归一化(原代码全局计算min/max不符合常规操作)。正确代码如下:
# 按列计算每个特征的min和max,保持维度以便广播 min_vals = conts_total.min(dim=0)[0].unsqueeze(0) max_vals = conts_total.max(dim=0)[0].unsqueeze(0) # 避免除以0(如果某列所有值相同) ranges = max_vals - min_vals ranges[ranges == 0] = 1.0 # 执行归一化 model_input = (conts_total - min_vals) / ranges
内容的提问来源于stack exchange,提问作者sarika
相关产品推荐
相关产品推荐

