已合并
[bugfix]跳过NZ格式Tensor的采集以避免异常内存申请 #795
jiangchao创建于 6月15日
[bugfix]跳过NZ格式Tensor的采集以避免异常内存申请 #795
已合并
jiangchao创建于 6月15日
共 3 个文件变更+26-20
@@ -366,12 +366,12 @@ dump.json is at ./dump_path/step*
366 * ①在非分布式场景下,如单进程训练或单卡训练中,训练进程没有`rank`信息,此时数据保存在`proc{pid}`,比对、分级可视化和溢出检测功能支持该目录下的数据解析。366 * ①在非分布式场景下,如单进程训练或单卡训练中,训练进程没有`rank`信息,此时数据保存在`proc{pid}`,比对、分级可视化和溢出检测功能支持该目录下的数据解析。
367 * ②在大模型训练过程中,可能既存在`rank`目录又存在`proc`目录,原因是一些进程可能仅在CPU上完成一些数据预处理操作,没有`rank`信息,此时目录名称为`proc{pid}`,这部分数据一般不存在精度问题,比对、分级可视化和溢出检测等功能将不会支持该目录下的数据解析。367 * ②在大模型训练过程中,可能既存在`rank`目录又存在`proc`目录,原因是一些进程可能仅在CPU上完成一些数据预处理操作,没有`rank`信息,此时目录名称为`proc{pid}`,这部分数据一般不存在精度问题,比对、分级可视化和溢出检测等功能将不会支持该目录下的数据解析。
368* `dump_tensor_data`:保存采集到的张量数据。368* `dump_tensor_data`:保存采集到的张量数据。
369-* `dump.json`:保存API或Module前反向数据的统计量信息。包含dump数据的API名称或Module名称,各数据的dtype、369+* `dump.json`:保存API或Module前反向数据的统计量信息。包含dump数据的API名称或Module名称,各数据的dtype、shape、max、min、mean、L2norm(L2范数,平方根)统计信息,以及根据`summary_mode`配置输出的校验值(`md5`对应CRC-32字段`md5`,`xor`对应XOR校验字段`md5`)。具体介绍可参考[dump.json文件说明](#dumpjson文件说明)。
370- shape、max、min、mean、L2norm(L2范数,平方根)统计信息,以及根据`summary_mode`370+ 
371- 配置输出的校验值(`md5`对应CRC-32字段`md5`,`xor`对应XOR校验字段`md5`)。371+ 当Tensor的数据格式为NZ格式时,dump.json中max、min、mean、L2norm(L2范数,平方根)统计信息均为`null`。且仅当Tensor的数据类型为torch.float32、torch.float16、torch.bfloat16中的一种时,才会计算mean、L2norm统计信息,其它数据类型仅计算max、min统计信息。
372- 具体介绍可参考[dump.json文件说明](#dumpjson文件说明)。
373 372 
374 当`summary_mode`配置为`xor`时,dump.json仅输出XOR校验值,不输出max、min、mean、L2norm统计信息。若安装的工具包编译时包含`--include-mod=xor_checksum`,PyTorch NPU场景会优先使用C++加速算子计算校验值,可带来数倍性能提升;安装方法请参见[安装基础工具包和xor_checksum加速算子](../msprobe_install_guide.md#install-xor-checksum)。加速算子不可用时自动回退到通用实现。373 当`summary_mode`配置为`xor`时,dump.json仅输出XOR校验值,不输出max、min、mean、L2norm统计信息。若安装的工具包编译时包含`--include-mod=xor_checksum`,PyTorch NPU场景会优先使用C++加速算子计算校验值,可带来数倍性能提升;安装方法请参见[安装基础工具包和xor_checksum加速算子](../msprobe_install_guide.md#install-xor-checksum)。加速算子不可用时自动回退到通用实现。
374+ 
375 当task配置为"nan_check"时,dump.json中各API数据将包含`is_nan`字段,取值为0或1,0代表无溢出状态,1代表有溢出状态(该模式下不保存API中数据的统计值)。375 当task配置为"nan_check"时,dump.json中各API数据将包含`is_nan`字段,取值为0或1,0代表无溢出状态,1代表有溢出状态(该模式下不保存API中数据的统计值)。
376* `dump_error_info.log`:仅在dump工具报错时生成此记录日志,用于记录dump错误日志。376* `dump_error_info.log`:仅在dump工具报错时生成此记录日志,用于记录dump错误日志。
377* `stack.json`:API/Module的调用栈信息。377* `stack.json`:API/Module的调用栈信息。
@@ -173,6 +173,13 @@ class TensorHandler:
173 if self.is_empty_data(common_tensor):173 if self.is_empty_data(common_tensor):
174 logger.debug(f"Saving fake tensor or meta tensor is not supported, the current tensor is {file_path}.")174 logger.debug(f"Saving fake tensor or meta tensor is not supported, the current tensor is {file_path}.")
175 return175 return
176+ if (
177+ common_tensor.device.type == "npu"
178+ and hasattr(torch_npu, "get_npu_format")
179+ and torch_npu.get_npu_format(common_tensor) == torch_npu.Format.FRACTAL_NZ
180+ ):
181+ logger.debug(f"Saving tensors with NZ format is not supported, the current tensor is {file_path}.")
182+ return
176 if common_tensor.untyped_storage().data_ptr() == 0:183 if common_tensor.untyped_storage().data_ptr() == 0:
177 logger.debug(f"Saving null-pointer tensor is not supported, the current tensor is {file_path}.")184 logger.debug(f"Saving null-pointer tensor is not supported, the current tensor is {file_path}.")
178 return185 return
@@ -372,6 +379,12 @@ class PytorchDataProcessor(BaseDataProcessor):
372 tensor_stat = TensorStatInfo()379 tensor_stat = TensorStatInfo()
373 if self.tensor_handler.is_empty_data(data) or self.tensor_handler.is_batchedtensor(data):380 if self.tensor_handler.is_empty_data(data) or self.tensor_handler.is_batchedtensor(data):
374 return tensor_stat381 return tensor_stat
382+ if (
383+ data.device.type == "npu"
384+ and hasattr(torch_npu, "get_npu_format")
385+ and torch_npu.get_npu_format(data) == torch_npu.Format.FRACTAL_NZ
386+ ):
387+ return tensor_stat
375 388 
376 data_clone = data.detach()389 data_clone = data.detach()
377 if self.tensor_handler.is_gradtrackingtensor(data_clone):390 if self.tensor_handler.is_gradtrackingtensor(data_clone):
@@ -405,22 +418,15 @@ class PytorchDataProcessor(BaseDataProcessor):
405 elif not data_clone.shape:418 elif not data_clone.shape:
406 tensor_stat.max = tensor_stat.min = tensor_stat.mean = tensor_stat.norm = data_clone.clone()419 tensor_stat.max = tensor_stat.min = tensor_stat.mean = tensor_stat.norm = data_clone.clone()
407 else:420 else:
408- if (
409- precision == Const.DUMP_PRECISION_HIGH
410- or data_clone.dtype == torch.float64
411- or not data_clone.is_floating_point()
412- ):
413- if data_clone.device.type == "npu" and data_clone.dtype == torch.int8:
414- try:
415- if torch_npu.get_npu_format(data_clone) == torch_npu.Format.FRACTAL_NZ:
416- return tensor_stat
417- except (AttributeError, RuntimeError):
418- return tensor_stat
419- data_clone = data_clone.float()
420 tensor_stat.max = torch.max(data_clone)421 tensor_stat.max = torch.max(data_clone)
421 tensor_stat.min = torch.min(data_clone)422 tensor_stat.min = torch.min(data_clone)
422- tensor_stat.mean = torch.mean(data_clone)423+ if (
423- tensor_stat.norm = torch.norm(data_clone)424+ precision != Const.DUMP_PRECISION_HIGH
425+ and data_clone.dtype != torch.float64
426+ and data_clone.is_floating_point()
427+ ):
428+ tensor_stat.mean = torch.mean(data_clone)
429+ tensor_stat.norm = torch.norm(data_clone)
424 return tensor_stat430 return tensor_stat
425 431 
426 def dump_async_data(self):432 def dump_async_data(self):
@@ -147,8 +147,8 @@ class TestPytorchDataProcessor(unittest.TestCase):
147 147 
148 self.assertEqual(result.max, 3)148 self.assertEqual(result.max, 3)
149 self.assertEqual(result.min, 1)149 self.assertEqual(result.min, 1)
150- self.assertEqual(result.mean, 2)150+ self.assertIsNone(result.mean)
151- self.assertEqual(result.norm, torch.norm(tensor.float()).item())151+ self.assertIsNone(result.norm)
152 152 
153 def test_get_stat_info_empty(self):153 def test_get_stat_info_empty(self):
154 tensor = torch.tensor([])154 tensor = torch.tensor([])