From be6e1d10ac78971b27159f154c42d6f99c7eb90a Mon Sep 17 00:00:00 2001 From: skydoorkai Date: Thu, 7 Aug 2025 13:30:30 +0800 Subject: [PATCH 1/4] update 250807 --- .../opt_lib/module_replace_optimization.py | 12 +- atorch/kernels/triton_jit/cross_entropy.py | 2 +- atorch/modules/fp8/cuda_kernel.py | 4 +- atorch/modules/fp8/quantize.py | 106 +- atorch/modules/fp8/triton_kernel.py | 31 +- atorch/ops/git_version_info_installed.py | 6 + atorch/pipeline_parallel/pipe_stage.py | 40 +- atorch/tests/common_tests/dim_planner_test.py | 113 -- .../tests/common_tests/dump_snapshot_test.py | 681 +++++++++++++ atorch/tests/common_tests/hebo_test.py | 39 +- atorch/tests/common_tests/inspector_test.py | 49 +- atorch/tests/common_tests/reorder_test.py | 75 ++ .../tests/common_tests/scaled_linear_test.py | 11 +- .../tests/common_tests/semi_auto_acc_test.py | 3 +- .../common_tests/test_dynamic_profile.py | 2 +- .../trainer/megatron_dataloader_test.py | 136 +++ .../{ => trainer}/trainer_test.py | 0 .../common_tests/trainer/trainer_v2_test.py | 747 ++++++++++++++ .../tests/common_tests/training_log_test.py | 350 +++++++ .../common_tests/unordered_dataloader_test.py | 3 + .../common_tests/virtual_optimizer_test.py | 335 ++++++ atorch/tests/toy_modules/toy_module_te.py | 61 ++ atorch/tests/utils/test_tools.py | 64 ++ atorch/trainer/args.py | 41 +- atorch/trainer/atorch_trainer_v2.py | 578 +++++------ atorch/trainer/debug_utils/debug_module.py | 132 +++ atorch/trainer/megatron/__init__.py | 4 +- .../trainer/megatron/megatron_async_save.py | 1 + .../trainer/megatron/megatron_ckpt_loader.py | 6 +- .../trainer/megatron/megatron_ckpt_saver.py | 19 + .../trainer/megatron/megatron_dataloader.py | 432 +++++--- atorch/trainer/megatron/megatron_wrapper.py | 961 +++++++++++++----- atorch/trainer/trainer_callback.py | 127 ++- atorch/trainer/utils.py | 92 +- atorch/utils/dynamic_profiler/__init__.py | 3 +- .../dynamic_profiler/_dynamic_profile.py | 596 +++++++++-- .../utils/dynamic_profiler/_file_monitor.py | 135 ++- atorch/utils/inspector/hooks.py | 128 ++- atorch/utils/parse_memory_pickle.py | 125 +++ atorch/utils/parse_trace_json.py | 134 ++- .../megatron_virtual_optimizer.py | 82 +- atorch/utils/virtual_optimizer/patch_utils.py | 176 ++++ atorch/utils/virtual_optimizer/pp_calc.py | 92 ++ .../baseline_megatron/pretrain_llama2_7b.sh | 0 .../{ => pretrain}/gpt2_config.yaml | 0 .../{run.sh => pretrain/launch_pretrain.sh} | 14 +- .../{ => pretrain}/llama2_7b_config.yaml | 28 +- .../pretrain_atorch_trainer_megatron.py | 24 +- examples/atorch_trainer_v2/sft/launch_sft.sh | 69 ++ .../sft/llama2_7b_config.yaml | 119 +++ .../sft/sft_atorch_trainer_megatron.py | 915 +++++++++++++++++ examples/moe/moe_modules.py | 25 +- 52 files changed, 6761 insertions(+), 1167 deletions(-) create mode 100644 atorch/ops/git_version_info_installed.py delete mode 100644 atorch/tests/common_tests/dim_planner_test.py create mode 100644 atorch/tests/common_tests/dump_snapshot_test.py create mode 100644 atorch/tests/common_tests/reorder_test.py create mode 100644 atorch/tests/common_tests/trainer/megatron_dataloader_test.py rename atorch/tests/common_tests/{ => trainer}/trainer_test.py (100%) create mode 100644 atorch/tests/common_tests/trainer/trainer_v2_test.py create mode 100644 atorch/tests/common_tests/training_log_test.py create mode 100644 atorch/tests/common_tests/virtual_optimizer_test.py create mode 100644 atorch/tests/toy_modules/toy_module_te.py create mode 100644 atorch/tests/utils/test_tools.py create mode 100644 atorch/trainer/debug_utils/debug_module.py create mode 100644 atorch/utils/parse_memory_pickle.py create mode 100644 atorch/utils/virtual_optimizer/patch_utils.py create mode 100644 atorch/utils/virtual_optimizer/pp_calc.py rename examples/atorch_trainer_v2/{ => pretrain}/baseline_megatron/pretrain_llama2_7b.sh (100%) rename examples/atorch_trainer_v2/{ => pretrain}/gpt2_config.yaml (100%) rename examples/atorch_trainer_v2/{run.sh => pretrain/launch_pretrain.sh} (82%) rename examples/atorch_trainer_v2/{ => pretrain}/llama2_7b_config.yaml (86%) rename examples/atorch_trainer_v2/{ => pretrain}/pretrain_atorch_trainer_megatron.py (98%) create mode 100755 examples/atorch_trainer_v2/sft/launch_sft.sh create mode 100644 examples/atorch_trainer_v2/sft/llama2_7b_config.yaml create mode 100644 examples/atorch_trainer_v2/sft/sft_atorch_trainer_megatron.py diff --git a/atorch/auto/opt_lib/module_replace_optimization.py b/atorch/auto/opt_lib/module_replace_optimization.py index f83563a..59df407 100644 --- a/atorch/auto/opt_lib/module_replace_optimization.py +++ b/atorch/auto/opt_lib/module_replace_optimization.py @@ -67,13 +67,15 @@ def decorator(cls): # decorator mode register. Not doing this in cls definition # because importing `register_replace_pair` there incurs circular import -register_replace_pair("HF_BertAttention_FA", supported_dtypes={torch.float16, torch.bfloat16})(BertAttentionFA) -register_replace_pair("HF_CLIPAttention_FA", supported_dtypes={torch.float16, torch.bfloat16})(CLIPAttentionFA) -register_replace_pair("MultiheadAttention_FA", supported_dtypes={torch.float16, torch.bfloat16})(MultiheadAttentionFA) if package_version_smaller_than("transformers", "4.38.0"): - # transformers 4.38.0 changed LlamaAttention interface, so check version first. + # Not support new transformer version, thus only apply to older version. register_replace_pair("HF_LlamaAttention_FA", supported_dtypes={torch.float16, torch.bfloat16})(LlamaAttentionFA) -register_replace_pair("HF_GPT2Attention_FA", supported_dtypes={torch.float16, torch.bfloat16})(GPT2AttentionFA) + register_replace_pair("HF_BertAttention_FA", supported_dtypes={torch.float16, torch.bfloat16})(BertAttentionFA) + register_replace_pair("HF_CLIPAttention_FA", supported_dtypes={torch.float16, torch.bfloat16})(CLIPAttentionFA) + register_replace_pair("MultiheadAttention_FA", supported_dtypes={torch.float16, torch.bfloat16})( + MultiheadAttentionFA + ) + register_replace_pair("HF_GPT2Attention_FA", supported_dtypes={torch.float16, torch.bfloat16})(GPT2AttentionFA) def _check_model_params_device(model): diff --git a/atorch/kernels/triton_jit/cross_entropy.py b/atorch/kernels/triton_jit/cross_entropy.py index 6106c26..c6314e5 100644 --- a/atorch/kernels/triton_jit/cross_entropy.py +++ b/atorch/kernels/triton_jit/cross_entropy.py @@ -64,7 +64,7 @@ def cross_entropy_fwd_kernel( else: label_idx -= class_start_idx if label_idx >= col_block_idx * BLOCK_SIZE and label_idx < min(n_cols, (col_block_idx + 1) * BLOCK_SIZE): - logits_label = tl.load(logits_ptr + label_idx) + logits_label = tl.load(logits_ptr + label_idx).to(tl.float32) # pragma: no cover if HAS_SMOOTHING: loss = ( (lse if not SPLIT else 0.0) diff --git a/atorch/modules/fp8/cuda_kernel.py b/atorch/modules/fp8/cuda_kernel.py index 64c4f8e..ce8c6df 100644 --- a/atorch/modules/fp8/cuda_kernel.py +++ b/atorch/modules/fp8/cuda_kernel.py @@ -20,11 +20,11 @@ def tile_quant( dtype=torch.float8_e4m3fn, block_size: int = 128, pow_2_scale: bool = False, + eps: float = 0.0, return_transpose: bool = False, use_cublas=False, ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor]]: # return qx, sx, qx_t, sx_t - eps = 1e-10 # only 128 is supported for block_size for now assert block_size == 128 return ops.quantize_vector_blockwise( @@ -43,11 +43,11 @@ def block_quant( dtype=torch.float8_e4m3fn, block_size: int = 128, pow_2_scale: bool = False, + eps: float = 0.0, return_transpose: bool = False, use_cublas=False, ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor]]: # return qx, sx, qx_t, sx_t - eps = 1e-10 # only 128 is supported for block_size for now assert block_size == 128 return ops.quantize_square_blockwise( diff --git a/atorch/modules/fp8/quantize.py b/atorch/modules/fp8/quantize.py index dff32d9..c6307fb 100644 --- a/atorch/modules/fp8/quantize.py +++ b/atorch/modules/fp8/quantize.py @@ -130,33 +130,45 @@ def get_linear_tileblock_quantize_params(block_size): class Fp8Quantization: E4M3_MAX_POS = torch.finfo(torch.float8_e4m3fn).max if hasattr(torch, "float8_e4m3fn") else 448.0 E5M2_MAX_POS = torch.finfo(torch.float8_e5m2).max if hasattr(torch, "float8_e5m2") else 57344.0 - EPS: float = 1e-12 @staticmethod - def _amax_to_scale(amax: torch.Tensor, float8_dtype: torch.dtype) -> torch.Tensor: + def _amax_to_scale(amax: torch.Tensor, float8_dtype: torch.dtype, eps: float = 0.0) -> torch.Tensor: with torch.no_grad(): - amax = amax.float() + amax = torch.clamp(amax.float(), min=eps) + if amax.numel() == 1: + if amax == 0.0: + amax.fill_(1.0) + else: + amax[amax == 0.0] = 1.0 + if float8_dtype == torch.float8_e4m3fn: - res = Fp8Quantization.E4M3_MAX_POS / torch.clamp(amax, min=Fp8Quantization.EPS) + res = Fp8Quantization.E4M3_MAX_POS / amax else: # e5m2 - res = Fp8Quantization.E5M2_MAX_POS / torch.clamp(amax, min=Fp8Quantization.EPS) + res = Fp8Quantization.E5M2_MAX_POS / amax return res @staticmethod def quantize_tensorwise( - x: torch.Tensor, float8_dtype: torch.dtype, method: ScaleComputMethod = ScaleComputMethod.DEFAULT + x: torch.Tensor, + float8_dtype: torch.dtype, + method: ScaleComputMethod = ScaleComputMethod.DEFAULT, + eps: float = 0.0, ): if method == ScaleComputMethod.DEFAULT or method == ScaleComputMethod.PYTORCH: - return Fp8Quantization.quantize_tensorwise_pt(x, float8_dtype) + return Fp8Quantization.quantize_tensorwise_pt(x, float8_dtype, eps) else: assert 0, f"{ScaleComputMethod} not implemented" @staticmethod def quantize_axiswise( - x: torch.Tensor, float8_dtype: torch.dtype, dim=1, method: ScaleComputMethod = ScaleComputMethod.DEFAULT + x: torch.Tensor, + float8_dtype: torch.dtype, + dim=1, + method: ScaleComputMethod = ScaleComputMethod.DEFAULT, + eps: float = 0.0, ): if method == ScaleComputMethod.DEFAULT or method == ScaleComputMethod.PYTORCH or x.dtype == torch.float16: - return Fp8Quantization.quantize_axiswise_pt(x, float8_dtype, dim) + return Fp8Quantization.quantize_axiswise_pt(x, float8_dtype, dim, eps) else: assert 0, f"{ScaleComputMethod} not implemented" @@ -167,10 +179,11 @@ def quantize_tilewise( block_size, method: ScaleComputMethod = ScaleComputMethod.DEFAULT, return_transpose=False, + eps=0.0, ): # Only CUTLASS kernel supports return_transpose if method == ScaleComputMethod.DEFAULT or method == ScaleComputMethod.TRITON: - return Fp8Quantization.quantize_tilewise_triton(x, float8_dtype, block_size=block_size) + return Fp8Quantization.quantize_tilewise_triton(x, float8_dtype, block_size=block_size, eps=eps) elif ( method == ScaleComputMethod.CUTLASS or method == ScaleComputMethod.CUBLAS @@ -180,6 +193,7 @@ def quantize_tilewise( x, float8_dtype, block_size=block_size, + eps=eps, return_transpose=return_transpose, use_cublas=(method != ScaleComputMethod.CUTLASS), ) @@ -193,6 +207,7 @@ def quantize_blockwise( block_size, method: ScaleComputMethod = ScaleComputMethod.DEFAULT, return_transpose=False, + eps=0.0, ): # Only CUTLASS kernel supports return_transpose if ( @@ -203,12 +218,13 @@ def quantize_blockwise( # Only support square block shape assert isinstance(block_size, int) or block_size[0] == block_size[1] bsize = block_size if isinstance(block_size, int) else block_size[0] - return Fp8Quantization.quantize_blockwise_triton(x, float8_dtype, block_size=bsize) + return Fp8Quantization.quantize_blockwise_triton(x, float8_dtype, block_size=bsize, eps=eps) elif method == ScaleComputMethod.CUTLASS or method == ScaleComputMethod.CUBLAS: return Fp8Quantization.quantize_blockwise_cuda( x, float8_dtype, block_size=block_size, + eps=eps, return_transpose=return_transpose, use_cublas=method == ScaleComputMethod.CUBLAS, ) @@ -216,52 +232,62 @@ def quantize_blockwise( assert 0, f"{ScaleComputMethod} not implemented" @staticmethod - def quantize_tensorwise_pt(x: torch.Tensor, float8_dtype: torch.dtype): + def quantize_tensorwise_pt(x: torch.Tensor, float8_dtype: torch.dtype, eps=0.0): amax = torch.max(torch.abs(x)) - scale = Fp8Quantization._amax_to_scale(amax, float8_dtype) + scale = Fp8Quantization._amax_to_scale(amax, float8_dtype, eps) x_fp8 = (x * scale).to(float8_dtype) inverse_scale = scale.reciprocal() return x_fp8, inverse_scale @staticmethod - def quantize_axiswise_pt(x: torch.Tensor, float8_dtype: torch.dtype, dim=1): + def quantize_axiswise_pt(x: torch.Tensor, float8_dtype: torch.dtype, dim=1, eps=0.0): # set dim=1 for rowwise, dim=0 for colwise. amax = torch.max(torch.abs(x), dim=dim, keepdim=True).values - scale = Fp8Quantization._amax_to_scale(amax, float8_dtype) + scale = Fp8Quantization._amax_to_scale(amax, float8_dtype, eps) x_fp8 = (x * scale).to(float8_dtype) inverse_scale = scale.reciprocal() return x_fp8, inverse_scale @staticmethod - def quantize_tilewise_triton(x: torch.Tensor, float8_dtype: torch.dtype, block_size=128): + def quantize_tilewise_triton(x: torch.Tensor, float8_dtype: torch.dtype, block_size=128, eps=0.0): from .triton_kernel import tile_quant - return tile_quant(x, dtype=float8_dtype, block_size=block_size) + return tile_quant(x, dtype=float8_dtype, block_size=block_size, eps=eps) @staticmethod def quantize_tilewise_cuda( - x: torch.Tensor, float8_dtype: torch.dtype, block_size=128, return_transpose=False, use_cublas=False + x: torch.Tensor, float8_dtype: torch.dtype, block_size=128, eps=0.0, return_transpose=False, use_cublas=False ): from .cuda_kernel import tile_quant return tile_quant( - x, dtype=float8_dtype, block_size=block_size, return_transpose=return_transpose, use_cublas=use_cublas + x, + dtype=float8_dtype, + block_size=block_size, + eps=eps, + return_transpose=return_transpose, + use_cublas=use_cublas, ) @staticmethod - def quantize_blockwise_triton(x: torch.Tensor, float8_dtype: torch.dtype, block_size=128): + def quantize_blockwise_triton(x: torch.Tensor, float8_dtype: torch.dtype, block_size=128, eps=0.0): from .triton_kernel import block_quant - return block_quant(x, dtype=float8_dtype, block_size=block_size) + return block_quant(x, dtype=float8_dtype, block_size=block_size, eps=eps) @staticmethod def quantize_blockwise_cuda( - x: torch.Tensor, float8_dtype: torch.dtype, block_size=128, return_transpose=False, use_cublas=False + x: torch.Tensor, float8_dtype: torch.dtype, block_size=128, eps=0.0, return_transpose=False, use_cublas=False ): from .cuda_kernel import block_quant return block_quant( - x, dtype=float8_dtype, block_size=block_size, return_transpose=return_transpose, use_cublas=use_cublas + x, + dtype=float8_dtype, + block_size=block_size, + eps=eps, + return_transpose=return_transpose, + use_cublas=use_cublas, ) @@ -276,6 +302,7 @@ def get_fp8_quantize_underflows( blockwise_required=True, block_size=128, quantize_method="DEFAULT", + eps=0.0, ): # Return a dict of underflow percentages for different quantize methods. if fp8_dtype is None: @@ -283,25 +310,32 @@ def get_fp8_quantize_underflows( data = data.contiguous().view(-1, data.shape[-1]) total_zero = (data == 0).sum().item() total_nonzero = data.numel() - total_zero + results = {} if tensorwise_required: - tensor_fp8, _ = Fp8Quantization.quantize_tensorwise(data, fp8_dtype) + tensor_fp8, _ = Fp8Quantization.quantize_tensorwise(data, fp8_dtype, eps=eps) tensor_zero = (tensor_fp8 == 0).sum().item() - tenor_underflow = (tensor_zero - total_zero) / total_nonzero * 100.0 + tenor_underflow = 0 + if total_nonzero > 0: + tenor_underflow = (tensor_zero - total_zero) / total_nonzero * 100.0 results["tensorwise"] = tenor_underflow if rowwise_required: row_fp8, _ = Fp8Quantization.quantize_axiswise( data, fp8_dtype, dim=1, method=ScaleComputMethod(quantize_method) ) row_zero = (row_fp8 == 0).sum().item() - row_underflow = (row_zero - total_zero) / total_nonzero * 100.0 + row_underflow = 0 + if total_nonzero > 0: + row_underflow = (row_zero - total_zero) / total_nonzero * 100.0 results["rowwise"] = row_underflow if colwise_required: col_fp8, _ = Fp8Quantization.quantize_axiswise( data, fp8_dtype, dim=0, method=ScaleComputMethod(quantize_method) ) col_zero = (col_fp8 == 0).sum().item() - col_underflow = (col_zero - total_zero) / total_nonzero * 100.0 + col_underflow = 0 + if total_nonzero > 0: + col_underflow = (col_zero - total_zero) / total_nonzero * 100.0 results["colwise"] = col_underflow padded_tensor = data @@ -312,25 +346,31 @@ def get_fp8_quantize_underflows( if tilewise_required: tile_fp8, _ = Fp8Quantization.quantize_tilewise( - padded_tensor, fp8_dtype, block_size, method=ScaleComputMethod(quantize_method) + padded_tensor, fp8_dtype, block_size, eps=eps, method=ScaleComputMethod(quantize_method) ) tile_zero = (tile_fp8 == 0).sum().item() - tile_underflow = (tile_zero - total_zero) / total_nonzero * 100.0 + tile_underflow = 0 + if total_nonzero > 0: + tile_underflow = (tile_zero - total_zero) / total_nonzero * 100.0 results["tilewise"] = tile_underflow if v_tilewise_required: tile_fp8, _ = Fp8Quantization.quantize_tilewise( - padded_tensor.t().contiguous(), fp8_dtype, block_size, method=ScaleComputMethod(quantize_method) + padded_tensor.t().contiguous(), fp8_dtype, block_size, eps=eps, method=ScaleComputMethod(quantize_method) ) tile_zero = (tile_fp8 == 0).sum().item() - tile_underflow = (tile_zero - total_zero) / total_nonzero * 100.0 + tile_underflow = 0 + if total_nonzero > 0: + tile_underflow = (tile_zero - total_zero) / total_nonzero * 100.0 results["v_tilewise"] = tile_underflow if blockwise_required: block_fp8, _ = Fp8Quantization.quantize_blockwise( - padded_tensor, fp8_dtype, block_size, method=ScaleComputMethod(quantize_method) + padded_tensor, fp8_dtype, block_size, eps=eps, method=ScaleComputMethod(quantize_method) ) block_zero = (block_fp8 == 0).sum().item() - block_underflow = (block_zero - total_zero) / total_nonzero * 100.0 + block_underflow = 0 + if total_nonzero > 0: + block_underflow = (block_zero - total_zero) / total_nonzero * 100.0 results["blockwise"] = block_underflow return results diff --git a/atorch/modules/fp8/triton_kernel.py b/atorch/modules/fp8/triton_kernel.py index 81bcb28..1d5e134 100644 --- a/atorch/modules/fp8/triton_kernel.py +++ b/atorch/modules/fp8/triton_kernel.py @@ -19,7 +19,7 @@ class RoundingMode(IntEnum): @triton.jit -def block_quant_kernel(x_ptr, y_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr): # pragma: no cover +def block_quant_kernel(x_ptr, y_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr, EPS: tl.constexpr): # pragma: no cover pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) n = tl.cdiv(N, BLOCK_SIZE) @@ -28,20 +28,21 @@ def block_quant_kernel(x_ptr, y_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr): # offs = offs_m[:, None] * N + offs_n[None, :] mask = (offs_m[:, None] < M) & (offs_n[None, :] < N) x = tl.load(x_ptr + offs, mask=mask).to(tl.float32) - eps = 1e-10 - s = tl.maximum(tl.max(tl.abs(x)), eps) / 448.0 + s = tl.maximum(tl.max(tl.abs(x)), EPS) / 448.0 + if s == 0.0: + s = 1.0 y = x / s y = y.to(y_ptr.dtype.element_ty) tl.store(y_ptr + offs, y, mask=mask) tl.store(s_ptr + pid_m * n + pid_n, s) -def block_quant(x: torch.Tensor, dtype=torch.float8_e4m3fn, block_size: int = 128) -> torch.Tensor: +def block_quant(x: torch.Tensor, dtype=torch.float8_e4m3fn, block_size: int = 128, eps: float = 0.0) -> torch.Tensor: M, N = x.size() y = torch.empty_like(x, dtype=dtype) s = x.new_empty(x.size(-2) // block_size, x.size(-1) // block_size, dtype=torch.float32) grid = lambda meta: (triton.cdiv(M, meta["BLOCK_SIZE"]), triton.cdiv(N, meta["BLOCK_SIZE"])) # noqa: E731 - block_quant_kernel[grid](x, y, s, M, N, BLOCK_SIZE=block_size) + block_quant_kernel[grid](x, y, s, M, N, BLOCK_SIZE=block_size, EPS=eps) return y, s @@ -92,13 +93,19 @@ def dequant(x: torch.Tensor, s: torch.Tensor, dtype=torch.float32, block_size: i @triton.jit def tile_quant_kernel( - x_ptr, y_ptr, s_ptr, scale_rounding_mode: tl.constexpr, BLOCK_SIZE: tl.constexpr + x_ptr, + y_ptr, + s_ptr, + scale_rounding_mode: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + EPS: tl.constexpr, ): # pragma: no cover pid = tl.program_id(axis=0) offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) x = tl.load(x_ptr + offs).to(tl.float32) - eps = 1e-10 - s = tl.maximum(tl.max(tl.abs(x)), eps) / 448.0 + s = tl.maximum(tl.max(tl.abs(x)), EPS) / 448.0 + if s == 0.0: + s = 1.0 if scale_rounding_mode == 1: # ceil rounding to power of 2 s = tl.ceil(tl.log2(s)) @@ -118,14 +125,18 @@ def tile_quant_kernel( def tile_quant( - x: torch.Tensor, dtype=torch.float8_e4m3fn, block_size: int = 128, scale_rounding_mode=RoundingMode.none + x: torch.Tensor, + dtype=torch.float8_e4m3fn, + block_size: int = 128, + scale_rounding_mode=RoundingMode.none, + eps: float = 0.0, ) -> Tuple[torch.Tensor, torch.Tensor]: assert x.is_contiguous() assert x.size(-1) % block_size == 0 y = torch.empty_like(x, dtype=dtype) s = x.new_empty(*x.size()[:-1], x.size(-1) // block_size, dtype=torch.float32) grid = lambda meta: (triton.cdiv(x.numel(), meta["BLOCK_SIZE"]),) # noqa: E731 - tile_quant_kernel[grid](x, y, s, scale_rounding_mode, BLOCK_SIZE=block_size) + tile_quant_kernel[grid](x, y, s, scale_rounding_mode, BLOCK_SIZE=block_size, EPS=eps) return y, s diff --git a/atorch/ops/git_version_info_installed.py b/atorch/ops/git_version_info_installed.py new file mode 100644 index 0000000..19189e5 --- /dev/null +++ b/atorch/ops/git_version_info_installed.py @@ -0,0 +1,6 @@ +version = "1.6.4dev8+45eb14116" +git_hash = "45eb14116" +git_branch = "1.6.4" +installed_ops = {"quantization_optimizer": False, "quantizer": False} +compatible_ops = {"quantization_optimizer": True, "quantizer": True, "atorch_not_implemented": False} +torch_info = {"version": "1.13", "bf16_support": False, "cuda_version": "11.6", "nccl_version": "2.14"} diff --git a/atorch/pipeline_parallel/pipe_stage.py b/atorch/pipeline_parallel/pipe_stage.py index e495445..e146545 100644 --- a/atorch/pipeline_parallel/pipe_stage.py +++ b/atorch/pipeline_parallel/pipe_stage.py @@ -1,20 +1,13 @@ +import functools import logging -from typing import Any, Callable, Dict, List, Optional, Tuple +from typing import Any, Callable, Dict, List, Optional, Tuple, cast import torch from torch import distributed as dist from torch import nn from atorch.communication.pipe_communicator import PipeCommunicator -from atorch.utils.version import torch_version - -logger = logging.getLogger(__name__) - -_is_torch_version_smaller_than_25 = False -if torch_version() < (2, 5, 0): # type: ignore - _is_torch_version_smaller_than_25 = True -else: - _is_torch_version_smaller_than_25 = False +from atorch.utils.version import torch_version as get_torch_version try: from torch.distributed.pipelining._debug import map_debug_info @@ -24,6 +17,9 @@ _PipelineStageBase = object map_debug_info, flatten_args = None, None +torch_version = cast(Tuple[int, ...], get_torch_version()) +logger = logging.getLogger(__name__) + class PipeStage(_PipelineStageBase): def __init__( @@ -61,7 +57,7 @@ def _create_grad_recv_info( ): pass - def _prepare_forward_infra(self, num_microbatches: int) -> None: + def _prepare_forward_infra(self, num_microbatches: int, *args, **kwargs) -> None: pass def _prepare_backward_infra(self, num_microbatches: int): @@ -222,7 +218,13 @@ def forward_one_chunk( return output - def backward_one_chunk(self, bwd_chunk_id: int, loss=None, full_backward: bool = True): + def backward_one_chunk( + self, + bwd_chunk_id: int, + loss=None, + full_backward: bool = True, + last_backward: Optional[bool] = None, + ): """ Perform backward pass on the module. This should only be called once per microbatch. @@ -255,24 +257,30 @@ def backward_one_chunk(self, bwd_chunk_id: int, loss=None, full_backward: bool = "input_values": input_values, } - if _is_torch_version_smaller_than_25: + if torch_version < (2, 5): assert full_backward, "only full_backward=True is supported" self.grads_input = self.backward_maybe_with_nosync(bwd_kwargs) else: assert self.dw_builder is None, "Not support dw builder" - + if torch_version < (2, 7): + backward_maybe_with_nosync = self.backward_maybe_with_nosync + else: + assert last_backward is not None + backward_maybe_with_nosync = functools.partial( + self.backward_maybe_with_nosync, last_backward=last_backward + ) # Save full_backward bwd_kwargs["full_backward"] = full_backward if full_backward: - self.grads_input, _ = self.backward_maybe_with_nosync("full", bwd_kwargs) + self.grads_input, _ = backward_maybe_with_nosync("full", bwd_kwargs) else: # perform the partial backwards for the inputs with a custom backward function # when the "stage_ouput" is a loss, then it is a tensor, otherwise it is a tuple of tensors if isinstance(bwd_kwargs["stage_output"], torch.Tensor): bwd_kwargs["stage_output"] = (bwd_kwargs["stage_output"],) - self.grads_input, param_groups = self.backward_maybe_with_nosync("input", bwd_kwargs) + self.grads_input, param_groups = backward_maybe_with_nosync("input", bwd_kwargs) # TODO: we dont need to save this, add to dw_runner? self.backward_state[bwd_chunk_id] = ( diff --git a/atorch/tests/common_tests/dim_planner_test.py b/atorch/tests/common_tests/dim_planner_test.py deleted file mode 100644 index 6cec591..0000000 --- a/atorch/tests/common_tests/dim_planner_test.py +++ /dev/null @@ -1,113 +0,0 @@ -import unittest - -import torch -from torch.nn import MSELoss -from torch.utils.data import Dataset - -from atorch.auto.model_context import ModelContext -from atorch.auto.opt_lib.shard_planners.dim_planner import DimPlanner -from atorch.utils.version import torch_version - -skip = False -try: - from pippy.IR import LossWrapper # noqa # type: ignore -except ImportError: - skip = True - - -class FakeDeviceContext(object): - """ - Device context contains compute resources information below. - number of nodes - number of logical cpu cores per node - gpu model - gpu memory(B) - number of gpus per node - total gpus of the training job - """ - - def __init__(self): - self.intra_node_bandwidth = 2**4 - self.inter_node_bandwidth = 2**3 - self.fp32_flops = 2**4 - self.gpu_memory = 2**13 - - -class MyModule(torch.nn.Module): - def __init__(self, in_features, out_features, bias=True): - super().__init__() - self.layer = torch.nn.Linear(in_features, out_features, bias=bias) - self.layers = torch.nn.ModuleList([torch.nn.Linear(out_features, out_features, bias=bias) for _ in range(3)]) - - def forward(self, input_): - data = torch.nn.functional.gelu(self.layer(input_[0])) - for op in self.layers: - data = op(data) - return data - - -class ToyDataset(Dataset): - def __init__(self, size): - self.size = size - - def __len__(self): - return self.size - - def __getitem__(self, idx): - return torch.ones((4, 8)), torch.zeros((4, 8)) - - -def prepare_input(data, device): - return data[0].to(device), data[1].to(device) - - -def my_loss_func(data, outputs): - loss_fct = MSELoss() - loss = loss_fct(outputs.view(-1), data[-1].view(-1)) - return loss - - -def create_model_context(data_size=512, batch_size=16, loss_func=None): - model = MyModule(8, 8, True) - dataset = ToyDataset(data_size) - model_context = ModelContext( - model=model, - optim_func=torch.optim.SGD, - dataset=dataset, - prepare_input=prepare_input, - dataloader_args={"batch_size": batch_size, "drop_last": True}, - optim_args={"lr": 0.001}, - loss_func=loss_func, - ) - return model_context - - -def run_dim_planner(num_nodes, num_devices_per_node, loss_func): - model_context = create_model_context(loss_func=loss_func) - device_context = FakeDeviceContext() - dim_planner = DimPlanner( - num_nodes=num_nodes, - num_devices_per_node=num_devices_per_node, - use_fake_mode=False, - device_context=device_context, - ) - ( - optimal_tensor_size, - optimal_pipe_size, - optimal_data_size, - insert_before_nodes, - ) = dim_planner.generate_sharding_plan(model_context) - assert len(insert_before_nodes) == optimal_pipe_size - 1 - assert optimal_tensor_size * optimal_pipe_size * optimal_data_size == num_nodes * num_devices_per_node - - -class TestDimPlanner(unittest.TestCase): - @unittest.skipIf( - not torch.cuda.is_available() or torch_version() < (2, 0, 0) or skip, "Test on GPU image" # type: ignore - ) - def test_dim_planner(self): - run_dim_planner(2, 4, my_loss_func) - - -if __name__ == "__main__": - unittest.main() diff --git a/atorch/tests/common_tests/dump_snapshot_test.py b/atorch/tests/common_tests/dump_snapshot_test.py new file mode 100644 index 0000000..79adc05 --- /dev/null +++ b/atorch/tests/common_tests/dump_snapshot_test.py @@ -0,0 +1,681 @@ +import glob +import json +import os +import subprocess +import sys +from functools import partial +from typing import Union +from urllib.request import urlretrieve + +import pytest +import torch + +import atorch +from atorch.common.log_utils import default_logger as logger + +pytestmark = pytest.mark.core24 + +pytest.importorskip("torch", minversion="2.0.9") +# if torch.version.git_version != "7bcf7da3a268b435777fe87c7794c382f444e86d" or not torch.cuda.is_available(): +# pytest.skip("requires pytorch 2.1 stable release", allow_module_level=True) + +try: + import ant_patches +except ModuleNotFoundError as e: + print(e) + print( + "Can't import ant_patches, if you want to use megatron with version >= 'core_r0.9.0', " + "please use 'ant_core_r0.9.0' branch." + ) + ant_patches = None + + +from atorch.common.util_func import find_free_port # noqa: E402 +from atorch.trainer.args import AtorchTrainingArgs # noqa: E402 +from atorch.trainer.atorch_trainer_v2 import AtorchTrainerV2 # noqa: E402 +from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 +from atorch.utils.import_util import is_coverage_available, is_megatron_lm_available # noqa: E402 +from atorch.utils.version import is_megatron_version_bigger_than, torch_version # noqa: E402 + +assert is_megatron_lm_available(), f"Can't import megatron, PYTHONPATH={os.environ['PYTHONPATH']}" + +if is_megatron_lm_available(): + import megatron.legacy.model + from megatron.core import mpu + from megatron.core.datasets.blended_megatron_dataset_builder import BlendedMegatronDatasetBuilder + from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset + from megatron.core.models.gpt import GPTModel + from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_local_spec, + get_gpt_layer_with_transformer_engine_spec, + ) + from megatron.core.transformer.spec_utils import import_module + from megatron.legacy.data.data_samplers import MegatronPretrainingRandomSampler, MegatronPretrainingSampler + from megatron.training import get_args, get_tokenizer, print_rank_0 + from megatron.training.arguments import core_transformer_config_from_args + from megatron.training.utils import ( + average_losses_across_data_parallel_group, + get_batch_on_this_cp_rank, + get_batch_on_this_tp_rank, + ) + from megatron.training.yaml_arguments import core_transformer_config_from_yaml + +if is_coverage_available(): + import coverage + + +def model_provider(pre_process=True, post_process=True) -> Union["GPTModel", "megatron.legacy.model.GPTModel"]: + """Builds the model. + + If you set the use_mcore_models to True, it will return the mcore GPT model and if not the legacy GPT model. + + Args: + pre_process (bool, optional): Set to true if you need to compute embedings. Defaults to True. + post_process (bool, optional): Set to true if you need to want to compute output logits/loss. Defaults to True. + + + Returns: + Union[GPTModel, megatron.legacy.model.GPTModel]: The returned model + """ + args = get_args() + use_te = args.transformer_impl == "transformer_engine" + + print_rank_0("building GPT model ...") + # Experimental loading arguments from yaml + if args.yaml_cfg is not None: + config = core_transformer_config_from_yaml(args, "language_model") + else: + config = core_transformer_config_from_args(args) + + if args.use_mcore_models: + if args.spec is not None: + transformer_layer_spec = import_module(args.spec) + else: + if use_te: + transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + args.num_experts, args.moe_grouped_gemm + ) + else: + transformer_layer_spec = get_gpt_layer_local_spec(args.num_experts, args.moe_grouped_gemm) + + model = GPTModel( + config=config, + transformer_layer_spec=transformer_layer_spec, + vocab_size=args.padded_vocab_size, + max_sequence_length=args.max_position_embeddings, + pre_process=pre_process, + post_process=post_process, + fp16_lm_cross_entropy=args.fp16_lm_cross_entropy, + parallel_output=True, + share_embeddings_and_output_weights=not args.untie_embeddings_and_output_weights, + position_embedding_type=args.position_embedding_type, + rotary_percent=args.rotary_percent, + ) + else: + assert args.context_parallel_size == 1, "Context parallelism is only supported with Megatron Core!" + + model = megatron.legacy.model.GPTModel( + config, num_tokentypes=0, parallel_output=True, pre_process=pre_process, post_process=post_process + ) + + return model + + +def train_valid_test_datasets_provider(train_val_test_num_samples): + """Build the train test and validation datasets. + + Args: + train_val_test_num_samples : A list containing the number of samples in train test and validation. + """ + + def is_dataset_built_on_rank(): + return ( + mpu.is_pipeline_first_stage() or mpu.is_pipeline_last_stage() + ) and mpu.get_tensor_model_parallel_rank() == 0 + + def core_gpt_dataset_config_from_args(args): + tokenizer = get_tokenizer() + + if is_megatron_version_bigger_than("0.6.0", check_equality=False): + from megatron.core.datasets.utils import get_blend_from_list + + return GPTDatasetConfig( + random_seed=args.seed, + sequence_length=args.seq_length, + blend=get_blend_from_list(args.data_path), + blend_per_split=[ + get_blend_from_list(args.train_data_path), + get_blend_from_list(args.valid_data_path), + get_blend_from_list(args.test_data_path), + ], + # renormalize_blend_weights=args.renormalize_blend_weights, + split=args.split, + num_dataset_builder_threads=args.num_dataset_builder_threads, + path_to_cache=args.data_cache_path, + mmap_bin_files=args.mmap_bin_files, + tokenizer=tokenizer, + reset_position_ids=args.reset_position_ids, + reset_attention_mask=args.reset_attention_mask, + eod_mask_loss=args.eod_mask_loss, + create_attention_mask=args.create_attention_mask_in_dataloader, + s3_cache_path=args.s3_cache_path, + ) + else: + return GPTDatasetConfig( + random_seed=args.seed, + sequence_length=args.seq_length, + blend=args.data_path, + blend_per_split=[ + args.train_data_path, + args.valid_data_path, + args.test_data_path, + ], + split=args.split, + path_to_cache=args.data_cache_path, + mock=args.mock_data, + mmap_bin_files=args.mmap_bin_files, + tokenizer=tokenizer, + reset_position_ids=args.reset_position_ids, + reset_attention_mask=args.reset_attention_mask, + eod_mask_loss=args.eod_mask_loss, + create_attention_mask=args.create_attention_mask_in_dataloader, + ) + + args = get_args() + config = core_gpt_dataset_config_from_args(args) + + if config.mock: + dataset_type = MockGPTDataset + else: + dataset_type = GPTDataset + + print_rank_0("> building train, validation, and test datasets for GPT ...") + + train_ds, valid_ds, test_ds = BlendedMegatronDatasetBuilder( + dataset_type, train_val_test_num_samples, is_dataset_built_on_rank, config + ).build() + + print_rank_0("> finished creating GPT datasets ...") + + return train_ds, valid_ds, test_ds + + +class DpoSampler(MegatronPretrainingSampler): + def set_epoch(self, epoch): + self.epoch = epoch + + +def build_train_valid_test_data_iterators(build_train_valid_test_datasets_provider): + """Build pretraining data iterators.""" + + def get_train_valid_test_num_samples(): + """Train/valid/test num samples.""" + + args = get_args() + + # Number of train/valid/test samples. + if args.train_samples: + train_samples = args.train_samples + else: + train_samples = args.train_iters * args.global_batch_size + eval_iters = (args.train_iters // args.eval_interval + 1) * args.eval_iters + if hasattr(args, "test_iters"): + test_iters = args.test_iters + else: + test_iters = args.eval_iters + + return ( + train_samples, + eval_iters * args.global_batch_size, + test_iters * args.global_batch_size, + ) + + def build_train_valid_test_datasets(build_train_valid_test_datasets_provider): + """Build pretraining datasets.""" + train_valid_test_num_samples = get_train_valid_test_num_samples() + print_rank_0(" > datasets target sizes (minimum size):") + print_rank_0(" train: {}".format(train_valid_test_num_samples[0])) + print_rank_0(" validation: {}".format(train_valid_test_num_samples[1])) + print_rank_0(" test: {}".format(train_valid_test_num_samples[2])) + return build_train_valid_test_datasets_provider(train_valid_test_num_samples) + + def build_pretraining_data_loader(dataset, consumed_samples): + """Build dataloader given an input dataset.""" + + if dataset is None: + return None + args = get_args() + + # Megatron sampler + if args.dataloader_type == "single": + # batch_sampler = DpoSampler( + batch_sampler = MegatronPretrainingSampler( + total_samples=len(dataset), + consumed_samples=consumed_samples, + micro_batch_size=args.micro_batch_size, + data_parallel_rank=mpu.get_data_parallel_rank(), + data_parallel_size=mpu.get_data_parallel_world_size(), + ) + elif args.dataloader_type == "cyclic": + batch_sampler = MegatronPretrainingRandomSampler( + dataset, + total_samples=len(dataset), + consumed_samples=consumed_samples, + micro_batch_size=args.micro_batch_size, + data_parallel_rank=mpu.get_data_parallel_rank(), + data_parallel_size=mpu.get_data_parallel_world_size(), + data_sharding=args.data_sharding, + ) + elif args.dataloader_type == "external": + # External dataloaders are passed through. User is expected to provide a + # torch-compatible dataloader and define samplers, if needed. + return dataset + else: + raise Exception("{} dataloader type is not supported.".format(args.dataloader_type)) + + # Torch dataloader. + return torch.utils.data.DataLoader( + dataset, + batch_sampler=batch_sampler, + num_workers=args.num_workers, + pin_memory=True, + persistent_workers=True if args.num_workers > 0 else False, + ) + + def build_train_valid_test_data_loaders(build_train_valid_test_datasets_provider): + """Build pretraining data loaders.""" + + args = get_args() + + (train_dataloader, valid_dataloader, test_dataloader) = (None, None, None) + + print_rank_0("> building train, validation, and test datasets ...") + + # For DPO ut + if not hasattr(args, "iteration"): + args.iteration = 0 + args.consumed_train_samples = 0 + + # Backward compatibility, assume fixed batch size. + if args.iteration > 0 and args.consumed_train_samples == 0: + assert args.train_samples is None, "only backward compatiblity support for iteration-based training" + args.consumed_train_samples = args.iteration * args.global_batch_size + if args.iteration > 0 and args.consumed_valid_samples == 0: + if args.train_samples is None: + args.consumed_valid_samples = ( + (args.iteration // args.eval_interval) * args.eval_iters * args.global_batch_size + ) + + # Rely on distributed-aware core datasets, temporary + is_distributed = getattr(build_train_valid_test_datasets_provider, "is_distributed", False) + + # Construct the data pipeline + if is_distributed or mpu.get_tensor_model_parallel_rank() == 0: + + # Build datasets. + train_ds, valid_ds, test_ds = build_train_valid_test_datasets(build_train_valid_test_datasets_provider) + # Build dataloders. + train_dataloader = build_pretraining_data_loader(train_ds, args.consumed_train_samples) + if args.skip_train: + valid_dataloader = build_pretraining_data_loader(valid_ds, 0) + else: + valid_dataloader = build_pretraining_data_loader(valid_ds, args.consumed_valid_samples) + test_dataloader = build_pretraining_data_loader(test_ds, 0) + + # Flags to know if we need to do training/validation/testing. + do_train = train_dataloader is not None and args.train_iters > 0 + do_valid = valid_dataloader is not None and args.eval_iters > 0 + do_test = test_dataloader is not None and args.eval_iters > 0 + flags = torch.tensor( + [int(do_train), int(do_valid), int(do_test)], + dtype=torch.long, + device="cuda", + ) + else: + flags = torch.tensor([0, 0, 0], dtype=torch.long, device="cuda") + + torch.distributed.broadcast(flags, 0) + + args.do_train = getattr(args, "do_train", False) or flags[0].item() + args.do_valid = getattr(args, "do_valid", False) or flags[1].item() + args.do_test = getattr(args, "do_test", False) or flags[2].item() + + return train_dataloader, valid_dataloader, test_dataloader + + args = get_args() + + # Build loaders. + train_dataloader, valid_dataloader, test_dataloader = build_train_valid_test_data_loaders( + build_train_valid_test_datasets_provider + ) + + if train_dataloader is not None: + logger.info(f"[Rank {args.rank}] build dataloader over!") + logger.info( + f"[Rank {args.rank}] train_dataloader {len(train_dataloader)} valid_dataloader {len(valid_dataloader)}" + f" test_dataloader {len(test_dataloader)}", + ) + logger.info( + f"[Rank {args.rank}] train_dataset {len(train_dataloader.dataset)}" + f" valid_dataset {len(valid_dataloader.dataset)} test_dataset {len(test_dataloader.dataset)}", + ) + + return train_dataloader, valid_dataloader, test_dataloader + + +class GPTTrainStep(MegatronTrainStep): + """ + GPT train step + + Args: + args (`argparse.Namespace`): Megatron-LM arguments. + """ + + def __init__(self, args, **kwargs): + super().__init__() + if not args.model_return_dict: + self.model_output_class = None + else: + from transformers.modeling_outputs import CausalLMOutputWithCrossAttentions + + self.model_output_class = CausalLMOutputWithCrossAttentions + + def get_batch_func(self, **kwargs): + def get_batch(data_iterator): + """Generate a batch.""" + + # TODO: this is pretty hacky, find a better way + if (not mpu.is_pipeline_first_stage()) and (not mpu.is_pipeline_last_stage()): + return None, None, None, None, None + + # get batches based on the TP rank you are on + batch = get_batch_on_this_tp_rank(data_iterator) + + # slice batch along sequence dimension for context parallelism + batch = get_batch_on_this_cp_rank(batch) + + return batch.values() + + return get_batch + + def get_loss_func(self, **kwargs): + def loss_func(loss_mask, output_tensor): + """Loss function. + + Args: + loss_mask (torch.Tensor): Used to mask out some portions of the loss + output_tensor (torch.Tensor): The tensor with the losses + """ + args = get_args() + + losses = output_tensor.float() + loss_mask = loss_mask.view(-1).float() + if args.context_parallel_size > 1: + loss = torch.cat([torch.sum(losses.view(-1) * loss_mask).view(1), loss_mask.sum().view(1)]) + torch.distributed.all_reduce(loss, group=mpu.get_context_parallel_group()) + loss = loss[0] / loss[1] + else: + loss = torch.sum(losses.view(-1) * loss_mask) / loss_mask.sum() + + # Check individual rank losses are not NaN prior to DP all-reduce. + if args.check_for_nan_in_loss_and_grad: + global_rank = torch.distributed.get_rank() + assert not loss.isnan(), ( + f"Rank {global_rank}: found NaN in local forward loss calculation. " + f"Device: {torch.cuda.current_device()}, node: {os.uname()[1]}" + ) + + # Reduce loss for logging. + averaged_loss = average_losses_across_data_parallel_group([loss]) + + return loss * args.context_parallel_size, {"lm loss": averaged_loss[0]} + + return loss_func + + def get_forward_step_func(self, **kwargs): + def forward_step(data_iterator, model: GPTModel): + """Forward training step. + + Args: + data_iterator : Input data iterator + model (GPTModel): The GPT Model + """ + # Get the batch. + tokens, labels, loss_mask, attention_mask, position_ids = self.get_batch_func()(data_iterator) + output_tensor = model(tokens, position_ids, attention_mask, labels=labels) + + return output_tensor, partial(self.get_loss_func(), loss_mask) + + return forward_step + + +vocab_url = "https://alps-common.oss-cn-hangzhou-zmf.aliyuncs.com/users/jinshi/atorch_unittest_data/vocab.json" +merge_url = "https://alps-common.oss-cn-hangzhou-zmf.aliyuncs.com/users/jinshi/atorch_unittest_data/merges.txt" +vocab_file = "/tmp/gpt2_vocab.json" +merge_file = "/tmp/gpt2_merges.json" + + +def download_tokenizer_file(): + try: + if not os.path.exists(vocab_file): + logger.info(f"Downloading {vocab_url} to {vocab_file}") + urlretrieve(vocab_url, vocab_file) + if not os.path.exists(merge_file): + logger.info(f"Downloading {merge_url} to {merge_file}") + urlretrieve(merge_url, merge_file) + except Exception as e: + logger.exception(f"Download {vocab_url} and {merge_url} failed, please check if the addresses exist. {e}") + return False + return True + + +def run_atorch_trainer_v2_dump_snapshot(): + output_dir = "/tmp/output_atorch_trainer" + + # test nv dynamic profiler + with open("/tmp/profile_config.json", "w") as f: + json.dump( + { + "mode": "dump", + "output_dir": "/tmp/profile", + "start_step": 2, + "schedule_warmup": 2, + "schedule_active": 1, + "with_stack": False, + "with_flops": False, + "with_modules": False, + "record_shapes": False, + "profile_memory": False, + "acc_events": False, + "activities": ["CPU", "CUDA"], + "profile_ranks": [-1], + "use_gzip": True, + "stacks": "all", + "context": "all", + "max_entries": 10000, + }, + f, + ) + + # test dynamic saving checkpoint + with open("/tmp/dynamic_save_config.json", "w") as f: + json.dump( + {"save_at_dynamic_steps": [100]}, + f, + ) + + training_args = AtorchTrainingArgs( + distributed_type="megatron", + output_dir=output_dir, + overwrite_output_dir=True, + per_device_train_batch_size=1, + per_device_eval_batch_size=1, + do_train=True, + bf16=True, + save_strategy="steps", + save_steps=20, + save_total_limit=1, + dynamic_save_config_path="/tmp/dynamic_save_config.json", + evaluation_strategy="steps", + eval_steps=25, + test_strategy="steps", + test_steps=25, + test_on_save=True, + logging_strategy="steps", + logging_steps=1, + logging_nan_inf_filter=False, + gradient_checkpointing=False, + tensorboard_dir=os.path.join(output_dir, "runs"), + use_deterministic_algorithms=True, + profiler_type="nv_dp", + dynamic_profiler_config_path="/tmp/profile_config.json", + memory_snapshot_path="/tmp/memory_snapshot", + # finetune_type="dpo", # to be removed + # max_steps=25, + ) + + train_valid_test_datasets_provider.is_distributed = True + + megatron_args = dict( + # Custom function + custom_model_provider_function=model_provider, + custom_megatron_dataloaders_provider_function=partial( + build_train_valid_test_data_iterators, train_valid_test_datasets_provider + ), + custom_train_step_class=GPTTrainStep, + # model args + model_type_name="gpt", + num_layers=16, + hidden_size=768, + num_attention_heads=12, + group_query_attention=True, + num_query_groups=12, + max_position_embeddings=512, + position_embedding_type="rope", + make_vocab_size_divisible_by=1, + norm_epsilon=1e-5, + normalization="RMSNorm", + untie_embeddings_and_output_weights=True, + use_flash_attn=True, + # tokenizer + tokenizer_type="GPT2BPETokenizer", + vocab_file=vocab_file, + merge_file=merge_file, + # optimizer + optimizer="adam", + # Regular args + attention_dropout=0.0, + hidden_dropout=0.0, + weight_decay=1e-1, + clip_grad=1.0, + adam_beta1=0.9, + adam_beta2=0.95, + adam_eps=1e-8, + # Megatron training args + pretraining_flag=True, + use_mcore_models=True, + transformer_impl="transformer_engine", + micro_batch_size=1, + global_batch_size=2, + add_bias_linear=False, + bias_gelu_fusion=False, + recompute_activations=True, + recompute_granularity="selective", + train_iters=8, + eval_iters=20, + test_iters=20, + overlap_grad_reduce=True, + overlap_param_gather=True, + # Distributed args + tensor_model_parallel_size=2, + pipeline_model_parallel_size=2, + num_virtual_stages_per_pipeline_rank=2, + sequence_parallel=True, + distributed_backend="nccl", + use_distributed_optimizer=True, + # Logging args + enable_one_logger=False, + log_timers_to_tensorboard=True, + log_validation_ppl_to_tensorboard=True, + log_memory_to_tensorboard=True, + log_throughput=True, + log_params_norm=True, + log_params_std=True, + tensorboard_dir=training_args.tensorboard_dir, + # Initialization args + seed=1403, + init_method_std=0.02, + # Learning rate args + lr=3e-5, + min_lr=3e-6, + lr_ecay_style="cosine", + lr_warmup_fraction=0.1, + # Data + data_cache_path=os.path.join(training_args.output_dir, "data_cache"), + mock_data=True, + seq_length=512, + num_workers=0, + mtp_num_layers=1, + ) + + if megatron_args["sequence_parallel"]: + os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "1" + + training_args.extra_configs = megatron_args + + trainer = AtorchTrainerV2( + args=training_args, + ) + train_result = trainer.train() + print(f"{train_result.metrics}") + + atorch.reset_distributed() + + +python_version = sys.version_info + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Skip cpu ut, only run on gpu.") +@pytest.mark.skipif(torch_version() < (2, 0, 0), reason="AtorchTrainer need torch2.0 .") # type: ignore +@pytest.mark.skipif( + not (python_version.major >= 3 and python_version.minor >= 10), reason="Megatron 0.11 requires python >= 3.10" +) +@pytest.mark.parametrize("gpu_num", [4]) +def test_atorch_trainer(gpu_num): + + if not download_tokenizer_file(): + logger.warning(f"Can't download {vocab_url} and {merge_url}, skip this unit test.") + return + + # Test for AntMonitor + if os.environ.get("ANTMONITOR_TFEVENT_PATH") is None: + os.environ["ANTMONITOR_TFEVENT_PATH"] = "/home/admin/logs/tfevent" + + launch_engine = "coverage run" if is_coverage_available() else "python" + dist_cmd = ( + f"{launch_engine} -m atorch.distributed.run --nnode=1 --nproc_per_node={gpu_num} " + f"--node_rank=0 --master_port={find_free_port()} {__file__}" + ) + + subprocess.run(dist_cmd, check=True, shell=True) + + # assert the gpu_num profile file is generated in /tmp/profile + profile_files = glob.glob("/tmp/profile/*.pickle") + assert len(profile_files) == gpu_num + + # assert the profile file is not empty + for profile_file in profile_files: + assert os.path.getsize(profile_file) > 0 + + +if __name__ == "__main__": + ut_cov = None + if is_coverage_available(): + ut_cov = coverage.Coverage() + ut_cov.start() + + run_atorch_trainer_v2_dump_snapshot() + + if ut_cov is not None: + ut_cov.stop() + ut_cov.save() diff --git a/atorch/tests/common_tests/hebo_test.py b/atorch/tests/common_tests/hebo_test.py index 24e657e..2358b92 100644 --- a/atorch/tests/common_tests/hebo_test.py +++ b/atorch/tests/common_tests/hebo_test.py @@ -5,34 +5,39 @@ import pandas as pd import torch -from atorch.auto.engine.sg_algo.hebo.acquisitions.acq import MACE, Acquisition, Mean, Sigma -from atorch.auto.engine.sg_algo.hebo.design_space.design_space import DesignSpace +try: + from atorch.auto.engine.sg_algo.hebo.acquisitions.acq import MACE, Acquisition, Mean, Sigma + from atorch.auto.engine.sg_algo.hebo.design_space.design_space import DesignSpace + class ToyExample(Acquisition): + def __init__(self, constr_v=1.0): + super().__init__(None) + self.constr_v = constr_v -class ToyExample(Acquisition): - def __init__(self, constr_v=1.0): - super().__init__(None) - self.constr_v = constr_v + @property + def num_obj(self): + return 1 - @property - def num_obj(self): - return 1 + @property + def num_constr(self): + return 1 - @property - def num_constr(self): - return 1 + def eval(self, x, xe): + # minimize L2norm(x) s.t. L2norm(x) > constr_v + out = (xe**2).sum(axis=1).reshape(-1, 1) + constr = self.constr_v - out + return np.concatenate([out, constr], axis=1) - def eval(self, x, xe): - # minimize L2norm(x) s.t. L2norm(x) > constr_v - out = (xe**2).sum(axis=1).reshape(-1, 1) - constr = self.constr_v - out - return np.concatenate([out, constr], axis=1) + DEP_AVAILABLE = True +except (ImportError, ModuleNotFoundError): + DEP_AVAILABLE = False def obj(x: pd.DataFrame) -> np.ndarray: return x["x0"].values.astype(float).reshape(-1, 1) ** 2 +@unittest.skipIf(not DEP_AVAILABLE, "missing dep") class HEBOTest(unittest.TestCase): def setUp(self): warnings.filterwarnings("ignore", category=DeprecationWarning) diff --git a/atorch/tests/common_tests/inspector_test.py b/atorch/tests/common_tests/inspector_test.py index b4b9374..c9f43ba 100644 --- a/atorch/tests/common_tests/inspector_test.py +++ b/atorch/tests/common_tests/inspector_test.py @@ -1,4 +1,5 @@ import unittest +from contextlib import nullcontext import pytest import torch @@ -7,6 +8,9 @@ from atorch.auto.accelerate import auto_accelerate from atorch.auto.opt_lib.amp_optimization import is_fp8_available from atorch.tests.toy_modules.toy_module import ToyDataset, ToyModel, loss_func, optim_func, prepare_input, run_train +from atorch.tests.toy_modules.toy_module_te import get_input as get_te_input +from atorch.tests.toy_modules.toy_module_te import get_model as get_te_model +from atorch.tests.toy_modules.toy_module_te import loss_func as te_loss_func from atorch.utils.inspector import TensorInspector try: @@ -17,6 +21,37 @@ _te_available = False +def get_fp8_context(is_init, enabled=True): + if not enabled: + return nullcontext() + fp8_format = transformer_engine.common.recipe.Format.E4M3 + fp8_recipe = transformer_engine.common.recipe.Float8BlockScaling(fp8_format=fp8_format) + if is_init: + context_args = {"enabled": True, "recipe": fp8_recipe} + + fp8_context = transformer_engine.pytorch.fp8_model_init(**context_args) + else: + fp8_context = transformer_engine.pytorch.fp8_autocast(enabled=True, fp8_recipe=fp8_recipe, fp8_group=None) + return fp8_context + + +def run_te_toy_model_with_inspector(hidden_size=256, num_gemms=4, use_fp8_init=False, use_fp8=False): + fp8_context_init = get_fp8_context(True, use_fp8_init) + with fp8_context_init: + model = get_te_model(hidden_size, num_gemms) + inspector = TensorInspector(2) + inspector.register_hooks(model, log_tensor_name_pattern="(gg)", te_fp8_check=use_fp8) + batch_size = 128 * num_gemms + for _ in range(5): + inputs = get_te_input(batch_size, hidden_size) + fp8_context = get_fp8_context(False, use_fp8) + with fp8_context: + outputs = model(inputs) + loss = te_loss_func(inputs, outputs) + loss.backward() + inspector.step() + + def run_toy_model_with_inspector( in_feature=16, out_feature=16, use_fp8=False, use_te=False, precision_switchable=False, summary_writer_items=None ): @@ -91,10 +126,11 @@ def run_toy_model_with_inspector( summary_writer_items=summary_writer_items, log_underflows_for_linear=True, ) - inspector.register_hooks(m_model, log_tensor_name_pattern="linears") + inspector.register_hooks(m_model, log_tensor_name_pattern="linears", te_fp8_check=True) device = "cuda" run_train(m_model, m_dataloader, m_optim, m_prepare_input, m_loss_func, device, inspector=inspector) + inspector.remove_hooks() summary_writer.close() @@ -126,3 +162,14 @@ def test_te_inspector(self): def test_scaled_linear_inspector(self): summary_writer_items = ["tensorwise_underflows", "rowwise_underflows"] run_toy_model_with_inspector(use_fp8=True, use_te=False, summary_writer_items=summary_writer_items) + + @unittest.skipIf( + not torch.cuda.is_available() or not is_fp8_available(), + "No fp8 gpu or te available for scaled linear tests", + ) + @pytest.mark.fp8 + def test_te_module_inspector(self): + # init fp8, compute fp8 + run_te_toy_model_with_inspector(use_fp8_init=True, use_fp8=True) + # init not fp8, compute fp8 + run_te_toy_model_with_inspector(use_fp8_init=False, use_fp8=True) diff --git a/atorch/tests/common_tests/reorder_test.py b/atorch/tests/common_tests/reorder_test.py new file mode 100644 index 0000000..66cb369 --- /dev/null +++ b/atorch/tests/common_tests/reorder_test.py @@ -0,0 +1,75 @@ +import os +import unittest + +from atorch.utils.rank_reorder.reorder import get_training_node_ranks_from_ips + + +class TestGetTrainingNodeRanksFromIps(unittest.TestCase): + def setUp(self): + os.environ["RANK_REORDER_WORKSPACE_DIR"] = "/tmp/" + os.environ["APP_ID"] = "test_reorder" + + def tearDown(self) -> None: + super().tearDown() + os.environ.pop("RANK_REORDER_WORKSPACE_DIR", None) + os.environ.pop("APP_ID", None) + + def test_empty_input(self): + result = get_training_node_ranks_from_ips([], {}) + self.assertEqual(result, {}) + + def test_all_nodes_paired(self): + training_node_ip_list = ["node1", "node2", "node3", "node4"] + cluster_tor_node_dict = {"tor1": ["node1", "node2"], "tor2": ["node3", "node4"]} + result = get_training_node_ranks_from_ips(training_node_ip_list, cluster_tor_node_dict) + expected_result = {"node1": 0, "node2": 1, "node3": 2, "node4": 3} + self.assertEqual(result, expected_result) + + def test_partially_paired_nodes(self): + training_node_ip_list = ["node1", "node2", "node3"] + cluster_tor_node_dict = {"tor1": ["node1"], "tor2": ["node2", "node3"]} + result = get_training_node_ranks_from_ips(training_node_ip_list, cluster_tor_node_dict) + expected_result = {"node1": 2, "node2": 0, "node3": 1} + self.assertEqual(result, expected_result) + + def test_three_nodes(self): + training_node_ip_list = ["node1", "node2", "node3", "node4"] + cluster_tor_node_dict = {"tor1": ["node1"], "tor2": ["node2", "node3", "node4"]} + result = get_training_node_ranks_from_ips(training_node_ip_list, cluster_tor_node_dict) + expected_result = {"node1": 0} + self.assertEqual(result, expected_result) + + def test_three_nodes_two_in_list(self): + training_node_ip_list = ["node1", "node2", "node3"] + cluster_tor_node_dict = {"tor1": ["node1"], "tor2": ["node2", "node3", "node4"]} + result = get_training_node_ranks_from_ips(training_node_ip_list, cluster_tor_node_dict) + expected_result = {"node1": 2, "node2": 0, "node3": 1} + self.assertEqual(result, expected_result) + + def test_empty_tor_dict(self): + training_node_ip_list = ["node1", "node2"] + cluster_tor_node_dict = {} + result = get_training_node_ranks_from_ips(training_node_ip_list, cluster_tor_node_dict) + self.assertEqual(result, {}) + + def test_multiple_tors(self): + training_node_ip_list = ["node1", "node2", "node3", "node4"] + cluster_tor_node_dict = { + "tor1": ["node1"], + "tor2": ["node2"], + "tor3": ["node3"], + "tor4": ["node4"], + } + result = get_training_node_ranks_from_ips(training_node_ip_list, cluster_tor_node_dict) + expected_result = {"node1": 0, "node2": 1, "node3": 2, "node4": 3} + self.assertEqual(result, expected_result) + + def test_nodes_not_in_cluster(self): + training_node_ip_list = ["node2", "node3"] + cluster_tor_node_dict = {"tor1": ["node1"]} + result = get_training_node_ranks_from_ips(training_node_ip_list, cluster_tor_node_dict) + self.assertEqual(result, {}) + + +if __name__ == "__main__": + unittest.main() diff --git a/atorch/tests/common_tests/scaled_linear_test.py b/atorch/tests/common_tests/scaled_linear_test.py index 345efa9..7c0ab8e 100644 --- a/atorch/tests/common_tests/scaled_linear_test.py +++ b/atorch/tests/common_tests/scaled_linear_test.py @@ -128,6 +128,9 @@ def test_quantize(self): data[0][0] = 10000.0 results = get_fp8_quantize_underflows(data) self.assertTrue("tilewise" in results and "blockwise" in results) + data = torch.zeros(128 * 4 - 3, 128 * 4, device="cuda", dtype=torch.bfloat16) + results = get_fp8_quantize_underflows(data) + self.assertTrue("tilewise" in results and "blockwise" in results) def switch_precision(self): layer_num = 6 @@ -251,6 +254,9 @@ def test_forward(self): model2 = ScaledLinear(K, N, bias=False, device="cuda", dtype=dtype, quantize_params=tileblock_triton_qparams) from atorch.modules.fp8.cuda_kernel import fp8_cutlass_cublas_ops_available + model3 = None + model4 = None + model5 = None if fp8_cutlass_cublas_ops_available(): model3 = ScaledLinear( K, N, bias=False, device="cuda", dtype=dtype, quantize_params=tileblock_cutlass_qparams @@ -263,11 +269,6 @@ def test_forward(self): model5 = ScaledLinear( K, N, bias=False, device="cuda", dtype=dtype, quantize_params=tileblock_deep_gemm_qparams ) - else: - model5 = None - else: - model3 = None - model4 = None model.set_use_fp8(False) with torch.no_grad(): diff --git a/atorch/tests/common_tests/semi_auto_acc_test.py b/atorch/tests/common_tests/semi_auto_acc_test.py index 6ed1a90..4f32c06 100644 --- a/atorch/tests/common_tests/semi_auto_acc_test.py +++ b/atorch/tests/common_tests/semi_auto_acc_test.py @@ -287,7 +287,8 @@ def run_gpt2_with_strategy( p_config = ([("data", zero_size)], None, True) else: p_config = None - if use_fa: + # new transformers version does not comparible with atorch fa + if use_fa and not use_fp8: strategy = [ ("parallel_mode", p_config), "module_replace", diff --git a/atorch/tests/common_tests/test_dynamic_profile.py b/atorch/tests/common_tests/test_dynamic_profile.py index ba619d0..9d717c8 100644 --- a/atorch/tests/common_tests/test_dynamic_profile.py +++ b/atorch/tests/common_tests/test_dynamic_profile.py @@ -67,7 +67,7 @@ class Bar: def test_thread_file_config_monitor(tmp_path): config_path = tmp_path / "test_config.json" monitor = ThreadFileConfigMonitor( - config_path.as_posix(), + [config_path.as_posix()], BarConfig, poll_interval=1, validator=lambda x: x.a == 1, diff --git a/atorch/tests/common_tests/trainer/megatron_dataloader_test.py b/atorch/tests/common_tests/trainer/megatron_dataloader_test.py new file mode 100644 index 0000000..9e94b29 --- /dev/null +++ b/atorch/tests/common_tests/trainer/megatron_dataloader_test.py @@ -0,0 +1,136 @@ +import os +import sys +from types import SimpleNamespace + +import pytest +import torch +import torch.distributed +import torch.multiprocessing as mp +from torch.utils.data import DataLoader, Dataset +from torch.utils.data.distributed import DistributedSampler + +import atorch +from atorch.common.util_func import find_free_port +from atorch.utils.import_util import is_megatron_lm_available, torch_version + +pytestmark = pytest.mark.core24 +pytest.importorskip("torch", minversion="2.0.9") + +python_version = sys.version_info + + +if is_megatron_lm_available(): + from megatron.core import parallel_state + from megatron.training.arguments import parse_args + from megatron.training.global_vars import set_args + + from atorch.trainer.megatron.megatron_dataloader import ( + MegatronDataloaderWrapper, + skip_first_batches_for_megatron_dataloader, + ) + + +class DummyDataset(Dataset): + def __init__(self, size=100, max_words=30): + self.size = size + self.max_words = max_words + + def __len__(self): + return self.size + + def __getitem__(self, idx): + input_ids = torch.ones([self.max_words], dtype=torch.int64) + labels = torch.ones([self.max_words], dtype=torch.int64) + attention_mask = torch.ones([self.max_words], dtype=torch.float32) + return { + "input_ids": input_ids, + "labels": labels, + "attention_mask": attention_mask, + } + + +def create_test_args(): + args = parse_args(ignore_unknown_args=True) + args.micro_batch_size = 2 + args.global_batch_size = 16 + args.data_parallel_size = atorch.world_size() + + return args + + +def _test_wrap_megatron_dataloader(rank, test_args): + os.environ["LOCAL_RANK"] = str(rank) + os.environ["RANK"] = str(rank) + backend = "nccl" if torch.cuda.is_available() else "gloo" + + if not atorch.init_distributed(backend, set_cuda_device_using_local_rank=True): + raise Exception("init failed") + + args = create_test_args() + set_args(args) + + args.rank = rank + args.world_size = atorch.world_size() + + parallel_state._MODEL_PARALLEL_GROUP = torch.distributed.new_group(backend=backend) + + if test_args.vpp_size == 0: + args.virtual_pipeline_model_parallel_size = None + dummy_dataset = DummyDataset(size=1000) + train_sampler = DistributedSampler(dummy_dataset, num_replicas=args.world_size, rank=args.rank) + dataloader = DataLoader(train_sampler, batch_size=args.micro_batch_size, sampler=train_sampler) + + megatron_dataloader = MegatronDataloaderWrapper(dataloader, is_post_training=True) + else: + args.virtual_pipeline_model_parallel_size = test_args.vpp_size + + dataloader = [None for _ in range(args.virtual_pipeline_model_parallel_size)] + + if args.rank in [0, args.world_size - 1]: + dummy_dataset = DummyDataset(size=1000) + train_sampler = DistributedSampler(dummy_dataset, num_replicas=args.world_size, rank=args.rank) + real_dataloader = DataLoader(train_sampler, batch_size=args.micro_batch_size, sampler=train_sampler) + + dataloader[0 if args.rank == 0 else args.virtual_pipeline_model_parallel_size - 1] = real_dataloader + + megatron_dataloader = MegatronDataloaderWrapper(dataloader, is_post_training=True) + + megatron_dataloader = skip_first_batches_for_megatron_dataloader(megatron_dataloader, num_batches=10) + + assert isinstance( + megatron_dataloader, MegatronDataloaderWrapper + ), f"'megatron_dataloader' should be MegatronDataloaderWrapper type, but got {type(megatron_dataloader)}." + + atorch.reset_distributed() + + +class TestMegatronDataloader: + @pytest.mark.skipif(not torch.cuda.is_available(), reason="Skip cpu ut, only run on gpu.") + @pytest.mark.skipif(torch_version() < (2, 0, 0), reason="AtorchTrainer need torch2.0 .") # type: ignore + @pytest.mark.skipif(torch.cuda.device_count() < 4, reason="run with cpu or gpu_num >=4") + @pytest.mark.skipif( + not (python_version.major >= 3 and python_version.minor >= 10), reason="Megatron 0.11 requires python >= 3.10" + ) + @pytest.mark.parametrize("vpp_size", [0, 4]) + def test_wrap_megatron_dataloader(self, vpp_size): + world_size = 4 + + os.environ["WORLD_SIZE"] = str(world_size) + os.environ["NPROC_PER_NODE"] = str(world_size) + os.environ["MASTER_ADDR"] = "localhost" + os.environ["MASTER_PORT"] = str(find_free_port()) + + test_args = SimpleNamespace() + test_args.vpp_size = vpp_size + + mp.spawn( + _test_wrap_megatron_dataloader, + args=(test_args,), + nprocs=world_size, + join=True, + daemon=False, + start_method="spawn", + ) + + os.environ["MASTER_ADDR"] = "" + os.environ["MASTER_PORT"] = "" diff --git a/atorch/tests/common_tests/trainer_test.py b/atorch/tests/common_tests/trainer/trainer_test.py similarity index 100% rename from atorch/tests/common_tests/trainer_test.py rename to atorch/tests/common_tests/trainer/trainer_test.py diff --git a/atorch/tests/common_tests/trainer/trainer_v2_test.py b/atorch/tests/common_tests/trainer/trainer_v2_test.py new file mode 100644 index 0000000..cbf0546 --- /dev/null +++ b/atorch/tests/common_tests/trainer/trainer_v2_test.py @@ -0,0 +1,747 @@ +import glob +import json +import os +import sys +from functools import partial +from typing import Union +from unittest.mock import MagicMock, patch +from urllib.request import urlretrieve + +import pytest +import torch +import torch.multiprocessing as mp + +import atorch +from atorch.common.log_utils import default_logger as logger + +pytestmark = pytest.mark.core24 +pytest.importorskip("torch", minversion="2.0.9") + +python_version = sys.version_info + +from atorch.common.util_func import find_free_port # noqa: E402 +from atorch.trainer.args import AtorchTrainingArgs # noqa: E402 +from atorch.trainer.atorch_trainer_v2 import AtorchTrainerV2 # noqa: E402 +from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 +from atorch.trainer.utils import DistributedType # noqa: E402 +from atorch.utils.import_util import is_megatron_lm_available # noqa: E402 +from atorch.utils.version import is_megatron_version_bigger_than, torch_version # noqa: E402 + +assert is_megatron_lm_available(), f"Can't import megatron, PYTHONPATH={os.environ['PYTHONPATH']}" + +if is_megatron_lm_available(): + import megatron.legacy.model + from megatron.core import mpu + from megatron.core.datasets.blended_megatron_dataset_builder import BlendedMegatronDatasetBuilder + from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset + from megatron.core.models.gpt import GPTModel + from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_local_spec, + get_gpt_layer_with_transformer_engine_spec, + ) + from megatron.core.transformer.spec_utils import import_module + from megatron.legacy.data.data_samplers import MegatronPretrainingRandomSampler, MegatronPretrainingSampler + from megatron.training import get_args, get_tokenizer, print_rank_0 + from megatron.training.arguments import core_transformer_config_from_args + from megatron.training.utils import ( + average_losses_across_data_parallel_group, + get_batch_on_this_cp_rank, + get_batch_on_this_tp_rank, + ) + from megatron.training.yaml_arguments import core_transformer_config_from_yaml + + +def model_provider(pre_process=True, post_process=True) -> Union["GPTModel", "megatron.legacy.model.GPTModel"]: + """Builds the model. + + If you set the use_mcore_models to True, it will return the mcore GPT model and if not the legacy GPT model. + + Args: + pre_process (bool, optional): Set to true if you need to compute embedings. Defaults to True. + post_process (bool, optional): Set to true if you need to want to compute output logits/loss. Defaults to True. + + + Returns: + Union[GPTModel, megatron.legacy.model.GPTModel]: The returned model + """ + args = get_args() + use_te = args.transformer_impl == "transformer_engine" + + print_rank_0("building GPT model ...") + # Experimental loading arguments from yaml + if args.yaml_cfg is not None: + config = core_transformer_config_from_yaml(args, "language_model") + else: + config = core_transformer_config_from_args(args) + + if args.use_mcore_models: + if args.spec is not None: + transformer_layer_spec = import_module(args.spec) + else: + if use_te: + transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + args.num_experts, args.moe_grouped_gemm + ) + else: + transformer_layer_spec = get_gpt_layer_local_spec(args.num_experts, args.moe_grouped_gemm) + + model = GPTModel( + config=config, + transformer_layer_spec=transformer_layer_spec, + vocab_size=args.padded_vocab_size, + max_sequence_length=args.max_position_embeddings, + pre_process=pre_process, + post_process=post_process, + fp16_lm_cross_entropy=args.fp16_lm_cross_entropy, + parallel_output=True, + share_embeddings_and_output_weights=not args.untie_embeddings_and_output_weights, + position_embedding_type=args.position_embedding_type, + rotary_percent=args.rotary_percent, + ) + else: + assert args.context_parallel_size == 1, "Context parallelism is only supported with Megatron Core!" + + model = megatron.legacy.model.GPTModel( + config, num_tokentypes=0, parallel_output=True, pre_process=pre_process, post_process=post_process + ) + + return model + + +def train_valid_test_datasets_provider(train_val_test_num_samples): + """Build the train test and validation datasets. + + Args: + train_val_test_num_samples : A list containing the number of samples in train test and validation. + """ + + def is_dataset_built_on_rank(): + return ( + mpu.is_pipeline_first_stage() or mpu.is_pipeline_last_stage() + ) and mpu.get_tensor_model_parallel_rank() == 0 + + def core_gpt_dataset_config_from_args(args): + tokenizer = get_tokenizer() + + if is_megatron_version_bigger_than("0.6.0", check_equality=False): + from megatron.core.datasets.utils import get_blend_from_list + + return GPTDatasetConfig( + random_seed=args.seed, + sequence_length=args.seq_length, + blend=get_blend_from_list(args.data_path), + blend_per_split=[ + get_blend_from_list(args.train_data_path), + get_blend_from_list(args.valid_data_path), + get_blend_from_list(args.test_data_path), + ], + # renormalize_blend_weights=args.renormalize_blend_weights, + split=args.split, + num_dataset_builder_threads=args.num_dataset_builder_threads, + path_to_cache=args.data_cache_path, + mmap_bin_files=args.mmap_bin_files, + tokenizer=tokenizer, + reset_position_ids=args.reset_position_ids, + reset_attention_mask=args.reset_attention_mask, + eod_mask_loss=args.eod_mask_loss, + create_attention_mask=args.create_attention_mask_in_dataloader, + s3_cache_path=args.s3_cache_path, + ) + else: + return GPTDatasetConfig( + random_seed=args.seed, + sequence_length=args.seq_length, + blend=args.data_path, + blend_per_split=[ + args.train_data_path, + args.valid_data_path, + args.test_data_path, + ], + split=args.split, + path_to_cache=args.data_cache_path, + mock=args.mock_data, + mmap_bin_files=args.mmap_bin_files, + tokenizer=tokenizer, + reset_position_ids=args.reset_position_ids, + reset_attention_mask=args.reset_attention_mask, + eod_mask_loss=args.eod_mask_loss, + create_attention_mask=args.create_attention_mask_in_dataloader, + ) + + args = get_args() + config = core_gpt_dataset_config_from_args(args) + + if config.mock: + dataset_type = MockGPTDataset + else: + dataset_type = GPTDataset + + print_rank_0("> building train, validation, and test datasets for GPT ...") + + train_ds, valid_ds, test_ds = BlendedMegatronDatasetBuilder( + dataset_type, train_val_test_num_samples, is_dataset_built_on_rank, config + ).build() + + print_rank_0("> finished creating GPT datasets ...") + + return train_ds, valid_ds, test_ds + + +class DpoSampler(MegatronPretrainingSampler): + def set_epoch(self, epoch): + self.epoch = epoch + + +def build_train_valid_test_data_iterators(build_train_valid_test_datasets_provider): + """Build pretraining data iterators.""" + + def get_train_valid_test_num_samples(): + """Train/valid/test num samples.""" + + args = get_args() + + # Number of train/valid/test samples. + if args.train_samples: + train_samples = args.train_samples + else: + train_samples = args.train_iters * args.global_batch_size + eval_iters = (args.train_iters // args.eval_interval + 1) * args.eval_iters + if hasattr(args, "test_iters"): + test_iters = args.test_iters + else: + test_iters = args.eval_iters + + return ( + train_samples, + eval_iters * args.global_batch_size, + test_iters * args.global_batch_size, + ) + + def build_train_valid_test_datasets(build_train_valid_test_datasets_provider): + """Build pretraining datasets.""" + train_valid_test_num_samples = get_train_valid_test_num_samples() + print_rank_0(" > datasets target sizes (minimum size):") + print_rank_0(" train: {}".format(train_valid_test_num_samples[0])) + print_rank_0(" validation: {}".format(train_valid_test_num_samples[1])) + print_rank_0(" test: {}".format(train_valid_test_num_samples[2])) + return build_train_valid_test_datasets_provider(train_valid_test_num_samples) + + def build_pretraining_data_loader(dataset, consumed_samples): + """Build dataloader given an input dataset.""" + + if dataset is None: + return None + args = get_args() + + # Megatron sampler + if args.dataloader_type == "single": + # batch_sampler = DpoSampler( + batch_sampler = MegatronPretrainingSampler( + total_samples=len(dataset), + consumed_samples=consumed_samples, + micro_batch_size=args.micro_batch_size, + data_parallel_rank=mpu.get_data_parallel_rank(), + data_parallel_size=mpu.get_data_parallel_world_size(), + ) + elif args.dataloader_type == "cyclic": + batch_sampler = MegatronPretrainingRandomSampler( + dataset, + total_samples=len(dataset), + consumed_samples=consumed_samples, + micro_batch_size=args.micro_batch_size, + data_parallel_rank=mpu.get_data_parallel_rank(), + data_parallel_size=mpu.get_data_parallel_world_size(), + data_sharding=args.data_sharding, + ) + elif args.dataloader_type == "external": + # External dataloaders are passed through. User is expected to provide a + # torch-compatible dataloader and define samplers, if needed. + return dataset + else: + raise Exception("{} dataloader type is not supported.".format(args.dataloader_type)) + + # Torch dataloader. + return torch.utils.data.DataLoader( + dataset, + batch_sampler=batch_sampler, + num_workers=args.num_workers, + pin_memory=True, + persistent_workers=True if args.num_workers > 0 else False, + ) + + def build_train_valid_test_data_loaders(build_train_valid_test_datasets_provider): + """Build pretraining data loaders.""" + + args = get_args() + + (train_dataloader, valid_dataloader, test_dataloader) = (None, None, None) + + print_rank_0("> building train, validation, and test datasets ...") + + # For DPO ut + if not hasattr(args, "iteration"): + args.iteration = 0 + args.consumed_train_samples = 0 + + # Backward compatibility, assume fixed batch size. + if args.iteration > 0 and args.consumed_train_samples == 0: + assert args.train_samples is None, "only backward compatiblity support for iteration-based training" + args.consumed_train_samples = args.iteration * args.global_batch_size + if args.iteration > 0 and args.consumed_valid_samples == 0: + if args.train_samples is None: + args.consumed_valid_samples = ( + (args.iteration // args.eval_interval) * args.eval_iters * args.global_batch_size + ) + + # Rely on distributed-aware core datasets, temporary + is_distributed = getattr(build_train_valid_test_datasets_provider, "is_distributed", False) + + # Construct the data pipeline + if is_distributed or mpu.get_tensor_model_parallel_rank() == 0: + + # Build datasets. + train_ds, valid_ds, test_ds = build_train_valid_test_datasets(build_train_valid_test_datasets_provider) + # Build dataloders. + train_dataloader = build_pretraining_data_loader(train_ds, args.consumed_train_samples) + if args.skip_train: + valid_dataloader = build_pretraining_data_loader(valid_ds, 0) + else: + valid_dataloader = build_pretraining_data_loader(valid_ds, args.consumed_valid_samples) + test_dataloader = build_pretraining_data_loader(test_ds, 0) + + # Flags to know if we need to do training/validation/testing. + do_train = train_dataloader is not None and args.train_iters > 0 + do_valid = valid_dataloader is not None and args.eval_iters > 0 + do_test = test_dataloader is not None and args.eval_iters > 0 + flags = torch.tensor( + [int(do_train), int(do_valid), int(do_test)], + dtype=torch.long, + device="cuda", + ) + else: + flags = torch.tensor([0, 0, 0], dtype=torch.long, device="cuda") + + torch.distributed.broadcast(flags, 0) + + args.do_train = getattr(args, "do_train", False) or flags[0].item() + args.do_valid = getattr(args, "do_valid", False) or flags[1].item() + args.do_test = getattr(args, "do_test", False) or flags[2].item() + + return train_dataloader, valid_dataloader, test_dataloader + + args = get_args() + + # Build loaders. + train_dataloader, valid_dataloader, test_dataloader = build_train_valid_test_data_loaders( + build_train_valid_test_datasets_provider + ) + + if train_dataloader is not None: + logger.info(f"[Rank {args.rank}] build dataloader over!") + logger.info( + f"[Rank {args.rank}] train_dataloader {len(train_dataloader)} valid_dataloader {len(valid_dataloader)}" + f" test_dataloader {len(test_dataloader)}", + ) + logger.info( + f"[Rank {args.rank}] train_dataset {len(train_dataloader.dataset)}" + f" valid_dataset {len(valid_dataloader.dataset)} test_dataset {len(test_dataloader.dataset)}", + ) + + return train_dataloader, valid_dataloader, test_dataloader + + +class GPTTrainStep(MegatronTrainStep): + """ + GPT train step + + Args: + args (`argparse.Namespace`): Megatron-LM arguments. + """ + + def __init__(self, args, **kwargs): + super().__init__() + if not args.model_return_dict: + self.model_output_class = None + else: + from transformers.modeling_outputs import CausalLMOutputWithCrossAttentions + + self.model_output_class = CausalLMOutputWithCrossAttentions + + def get_batch_func(self, **kwargs): + def get_batch(data_iterator): + """Generate a batch.""" + + # TODO: this is pretty hacky, find a better way + if (not mpu.is_pipeline_first_stage()) and (not mpu.is_pipeline_last_stage()): + return None, None, None, None, None + + # get batches based on the TP rank you are on + batch = get_batch_on_this_tp_rank(data_iterator) + + # slice batch along sequence dimension for context parallelism + batch = get_batch_on_this_cp_rank(batch) + + return batch.values() + + return get_batch + + def get_loss_func(self, **kwargs): + def loss_func(loss_mask, output_tensor): + """Loss function. + + Args: + loss_mask (torch.Tensor): Used to mask out some portions of the loss + output_tensor (torch.Tensor): The tensor with the losses + """ + args = get_args() + + losses = output_tensor.float() + loss_mask = loss_mask.view(-1).float() + if args.context_parallel_size > 1: + loss = torch.cat([torch.sum(losses.view(-1) * loss_mask).view(1), loss_mask.sum().view(1)]) + torch.distributed.all_reduce(loss, group=mpu.get_context_parallel_group()) + loss = loss[0] / loss[1] + else: + loss = torch.sum(losses.view(-1) * loss_mask) / loss_mask.sum() + + # Check individual rank losses are not NaN prior to DP all-reduce. + if args.check_for_nan_in_loss_and_grad: + global_rank = torch.distributed.get_rank() + assert not loss.isnan(), ( + f"Rank {global_rank}: found NaN in local forward loss calculation. " + f"Device: {torch.cuda.current_device()}, node: {os.uname()[1]}" + ) + + # Reduce loss for logging. + averaged_loss = average_losses_across_data_parallel_group([loss]) + + return loss * args.context_parallel_size, {"lm loss": averaged_loss[0]} + + return loss_func + + def get_forward_step_func(self, **kwargs): + def forward_step(data_iterator, model: GPTModel): + """Forward training step. + + Args: + data_iterator : Input data iterator + model (GPTModel): The GPT Model + """ + # Get the batch. + tokens, labels, loss_mask, attention_mask, position_ids = self.get_batch_func()(data_iterator) + output_tensor = model(tokens, position_ids, attention_mask, labels=labels) + + return output_tensor, partial(self.get_loss_func(), loss_mask) + + return forward_step + + +vocab_url = "https://alps-common.oss-cn-hangzhou-zmf.aliyuncs.com/users/jinshi/atorch_unittest_data/vocab.json" +merge_url = "https://alps-common.oss-cn-hangzhou-zmf.aliyuncs.com/users/jinshi/atorch_unittest_data/merges.txt" +vocab_file = "/tmp/gpt2_vocab.json" +merge_file = "/tmp/gpt2_merges.json" + + +def download_tokenizer_file(): + try: + if not os.path.exists(vocab_file): + logger.info(f"Downloading {vocab_url} to {vocab_file}") + urlretrieve(vocab_url, vocab_file) + if not os.path.exists(merge_file): + logger.info(f"Downloading {merge_url} to {merge_file}") + urlretrieve(merge_url, merge_file) + except Exception as e: + logger.exception(f"Download {vocab_url} and {merge_url} failed, please check if the addresses exist. {e}") + return False + return True + + +def run_atorch_trainer_v2(rank): + os.environ["LOCAL_RANK"] = str(rank) + os.environ["RANK"] = str(rank) + + output_dir = "/tmp/output_atorch_trainer" + + # test nv dynamic profiler + with open("/tmp/profile_config.json", "w") as f: + json.dump( + { + "output_dir": "/tmp/profile", + "start_step": 20, + "schedule_warmup": 2, + "schedule_active": 1, + "with_stack": False, + "with_flops": False, + "with_modules": False, + "record_shapes": False, + "profile_memory": False, + "acc_events": False, + "activities": ["CPU", "CUDA"], + "profile_ranks": [-1], + "use_gzip": True, + }, + f, + ) + + # test dynamic saving checkpoint + with open("/tmp/dynamic_save_config.json", "w") as f: + json.dump( + {"save_at_dynamic_steps": [100]}, + f, + ) + + training_args = AtorchTrainingArgs( + distributed_type="megatron", + output_dir=output_dir, + overwrite_output_dir=True, + per_device_train_batch_size=1, + per_device_eval_batch_size=1, + do_train=True, + bf16=True, + save_strategy="steps", + save_steps=20, + save_total_limit=1, + dynamic_save_config_path="/tmp/dynamic_save_config.json", + evaluation_strategy="steps", + eval_steps=25, + test_strategy="steps", + test_steps=25, + test_on_save=True, + logging_strategy="steps", + logging_steps=1, + logging_nan_inf_filter=False, + gradient_checkpointing=False, + tensorboard_dir=os.path.join(output_dir, "runs"), + use_deterministic_algorithms=True, + profiler_type="nv_dp", + dynamic_profiler_config_path="/tmp/profile_config.json", + memory_snapshot_path="/tmp/memory_snapshot", + # finetune_type="dpo", # to be removed + # max_steps=25, + ) + + train_valid_test_datasets_provider.is_distributed = True + + megatron_args = dict( + # Custom function + custom_model_provider_function=model_provider, + custom_megatron_dataloaders_provider_function=partial( + build_train_valid_test_data_iterators, train_valid_test_datasets_provider + ), + custom_train_step_class=GPTTrainStep, + # model args + model_type_name="gpt", + num_layers=16, + hidden_size=768, + num_attention_heads=12, + group_query_attention=True, + num_query_groups=12, + max_position_embeddings=512, + position_embedding_type="rope", + make_vocab_size_divisible_by=1, + norm_epsilon=1e-5, + normalization="RMSNorm", + untie_embeddings_and_output_weights=True, + use_flash_attn=True, + # tokenizer + tokenizer_type="GPT2BPETokenizer", + vocab_file=vocab_file, + merge_file=merge_file, + # optimizer + optimizer="adam", + # Regular args + attention_dropout=0.0, + hidden_dropout=0.0, + weight_decay=1e-1, + clip_grad=1.0, + adam_beta1=0.9, + adam_beta2=0.95, + adam_eps=1e-8, + # Megatron training args + pretraining_flag=True, + use_mcore_models=True, + transformer_impl="transformer_engine", + micro_batch_size=1, + global_batch_size=2, + add_bias_linear=False, + bias_gelu_fusion=False, + recompute_activations=True, + recompute_granularity="selective", + train_iters=25, + eval_iters=5, + test_iters=5, + overlap_grad_reduce=True, + overlap_param_gather=True, + # Distributed args + tensor_model_parallel_size=2, + pipeline_model_parallel_size=2, + num_virtual_stages_per_pipeline_rank=2, + sequence_parallel=True, + distributed_backend="nccl", + use_distributed_optimizer=True, + # Logging args + enable_one_logger=False, + log_timers_to_tensorboard=True, + log_validation_ppl_to_tensorboard=True, + log_memory_to_tensorboard=True, + log_throughput=True, + log_params_norm=True, + log_params_std=True, + tensorboard_dir=training_args.tensorboard_dir, + # Initialization args + seed=1403, + init_method_std=0.02, + # Learning rate args + lr=3e-5, + min_lr=3e-6, + lr_ecay_style="cosine", + lr_warmup_fraction=0.1, + # Data + data_cache_path=os.path.join(training_args.output_dir, "data_cache"), + mock_data=True, + seq_length=512, + num_workers=0, + mtp_num_layers=1, + # routing_map save + moe_router_save=True, + moe_token_dispatcher_type="alltoall", + moe_router_save_dir="/tmp/trysave", + moe_splits_save_dir="/tmp/splitsave/", + moe_router_save_iters="2,5,10,15", + moe_router_load=True, + moe_router_load_dir="/tmp/trysave/iter_10/", + ) + + if megatron_args["sequence_parallel"]: + os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "1" + + training_args.extra_configs = megatron_args + example_data = {"routing_map": torch.randn(10, 10), "metadata": {"dp_rank": 1, "layer_id": 0}} + os.makedirs("/tmp/trysave/iter_10/", exist_ok=True) + filename = "layer_expert_layer_id_0_dp_rank_1_routing_map.pt" + torch.save(example_data, os.path.join("/tmp/trysave/iter_10/", filename)) + trainer = AtorchTrainerV2( + args=training_args, + ) + train_result = trainer.train() + print(f"{train_result.metrics}") + + atorch.reset_distributed() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Skip cpu ut, only run on gpu.") +@pytest.mark.skipif(torch_version() < (2, 0, 0), reason="AtorchTrainer need torch2.0 .") # type: ignore +@pytest.mark.skipif( + not (python_version.major >= 3 and python_version.minor >= 10), reason="Megatron 0.11 requires python >= 3.10" +) +@pytest.mark.parametrize("world_size", [4]) +def test_atorch_trainer(world_size): + + if not download_tokenizer_file(): + logger.warning(f"Can't download {vocab_url} and {merge_url}, skip this unit test.") + return + + # Test for AntMonitor + if os.environ.get("ANTMONITOR_TFEVENT_PATH") is None: + os.environ["ANTMONITOR_TFEVENT_PATH"] = "/home/admin/logs/tfevent" + + os.environ["WORLD_SIZE"] = str(world_size) + os.environ["NPROC_PER_NODE"] = str(world_size) + os.environ["MASTER_ADDR"] = "localhost" + os.environ["MASTER_PORT"] = str(find_free_port()) + + mp.spawn( + run_atorch_trainer_v2, + nprocs=world_size, + join=True, + daemon=False, + start_method="spawn", + ) + + os.environ["MASTER_ADDR"] = "" + os.environ["MASTER_PORT"] = "" + + # assert the gpu_num profile file is generated in /tmp/profile + profile_files = glob.glob("/tmp/profile/*.json.gz") + assert len(profile_files) == world_size + + # assert the profile file is not empty + for profile_file in profile_files: + assert os.path.getsize(profile_file) > 0 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Skip cpu ut, only run on gpu.") +@pytest.mark.skipif(torch_version() < (2, 0, 0), reason="AtorchTrainer need torch2.0 .") # type: ignore +@pytest.mark.skipif( + not (python_version.major >= 3 and python_version.minor >= 10), reason="Megatron 0.11 requires python >= 3.10" +) +@pytest.mark.parametrize("distributed_type", ["megatron"]) +def test_evaluate_with_mocks(distributed_type): + with patch("atorch.trainer.args.AtorchAcceleratorState") as mock_accel_state, patch( + "atorch.trainer.atorch_trainer_v2.get_timers" + ) as mock_get_timers, patch("atorch.trainer.atorch_trainer_v2.get_args") as mock_get_args: + + # Patch AtorchAcceleratorState to have required attributes + mock_accel_state.return_value.distributed_type = DistributedType.MEGATRON + mock_accel_state.return_value.is_main_process = True + mock_accel_state.return_value.is_local_main_process = True + mock_accel_state.return_value.local_process_index = 0 + + # Patch timers + mock_timer = MagicMock() + mock_get_timers.return_value = MagicMock(return_value=mock_timer) + mock_timer.start.return_value = None + mock_timer.stop.return_value = None + mock_timer.log.return_value = None + + # Patch megatron args + mock_args = MagicMock() + mock_args.test_iters = 0 + mock_args.global_batch_size = 1 + mock_get_args.return_value = mock_args + + # Import after patching AtorchAcceleratorState + from atorch.trainer.args import AtorchTrainingArgs + from atorch.trainer.atorch_trainer_v2 import AtorchTrainerV2 + + training_args = AtorchTrainingArgs( + distributed_type=distributed_type, + output_dir="/tmp/test_eval", + overwrite_output_dir=True, + per_device_eval_batch_size=1, + do_eval=True, + ) + + # Patch train_engine and its methods + trainer = AtorchTrainerV2(args=training_args) + + # Patch train_engine and its methods + trainer.train_engine = MagicMock() + trainer.train_engine.get_dataloader.return_value = [torch.tensor(1.0), torch.tensor(2.0)] + trainer.train_engine.eval.return_value = None + trainer.train_engine.train.return_value = None + trainer.train_engine.train_step_handler.model_output_class = None + + # Patch callback_handler + trainer.callback_handler = MagicMock() + trainer.callback_handler.on_evaluate_begin.return_value = None + trainer.callback_handler.on_prediction_step.return_value = None + trainer.callback_handler.on_evaluate.return_value = None + + # Patch log + trainer.log = MagicMock() + + # Patch control/state + trainer.control = MagicMock() + trainer.state = MagicMock() + + # Run evaluate + result = trainer.evaluate(eval_or_test="test") + print(result) + assert isinstance(result, dict) + assert "test_loss" in result or "test_perplexity" in result + + trainer.train_engine.get_dataloader.return_value = None + result = trainer.evaluate(eval_or_test="test") + print(result) + assert result is None diff --git a/atorch/tests/common_tests/training_log_test.py b/atorch/tests/common_tests/training_log_test.py new file mode 100644 index 0000000..35177a0 --- /dev/null +++ b/atorch/tests/common_tests/training_log_test.py @@ -0,0 +1,350 @@ +from unittest.mock import MagicMock, patch + +import pytest +import torch + +from atorch.utils.import_util import is_megatron_lm_available + +if is_megatron_lm_available(): + from megatron.core import utils as megatron_utils + from megatron.legacy import model as megatron_legacy_model + + from atorch.trainer.utils import training_log +else: + megatron_utils = MagicMock() + megatron_legacy_model = MagicMock() + + +MOCK_BASE = "atorch.trainer.utils" + + +def _setup_training_log_mocks( + mock_datetime, + mock_mpu, + mock_report_theoretical_memory, + mock_report_memory, + mock_num_fp_ops, + mock_mem_stats, + mock_version_check, + mock_print_rank_last, + mock_get_num_microbatches, + mock_get_one_logger, + mock_get_wandb_writer, + mock_get_tensorboard_writer, + mock_get_timers, + mock_get_args, +): + """Helper to setup common mocks for training_log tests.""" + # Setup mocks + mock_args = MagicMock() + mock_get_args.return_value = mock_args + + mock_timers_instance = MagicMock() + mock_timers_instance.log = MagicMock() + interval_timer = MagicMock() + interval_timer.elapsed.return_value = 1.0 + mock_timers_instance.return_value = interval_timer + mock_get_timers.return_value = mock_timers_instance + + mock_writer = MagicMock() + mock_get_tensorboard_writer.return_value = mock_writer + + mock_datetime.now.return_value.strftime.return_value = "2023-01-01 00:00:00" + + # Common args + mock_args.log_interval = 10 + mock_args.use_local_sgd = False + mock_args.num_experts = None + mock_args.mtp_num_layers = None + mock_args.log_throughput = False + mock_args.decoupled_lr = None + mock_args.data_parallel_size = 1 + mock_args.micro_batch_size = 1 + mock_args.world_size = 1 + mock_args.consumed_train_samples = 100 + mock_args.train_iters = 100 + mock_args.tensorboard_log_interval = 5 + mock_args.skipped_train_samples = 0 + mock_args.log_memory_to_tensorboard = False + + mock_mpu.is_pipeline_first_stage.return_value = True + mock_mpu.is_pipeline_last_stage.return_value = True + + # Dummy inputs for training_log + loss_dict = {"lm_loss": torch.tensor(1.0)} + total_loss_dict = {} + common_kwargs = dict( + loss_dict=loss_dict, + total_loss_dict=total_loss_dict, + learning_rate=0.001, + decoupled_learning_rate=None, + loss_scale=1.0, + report_memory_flag=False, + skipped_iter=0, + grad_norm=None, + params_norm=None, + num_zeros_in_grad=None, + params_std=None, + custom_metrics=None, + ) + + # To enter the main logging block, `iteration % args.log_interval` must be 0. + # We set `log_interval` to 10. + iteration = 10 + + return mock_args, mock_timers_instance, common_kwargs, iteration + + +@pytest.mark.skipif(not is_megatron_lm_available(), reason="Megatron-LM not available.") +@patch(f"{MOCK_BASE}.get_args") +@patch(f"{MOCK_BASE}.get_timers") +@patch(f"{MOCK_BASE}.get_tensorboard_writer") +@patch(f"{MOCK_BASE}.get_wandb_writer") +@patch(f"{MOCK_BASE}.get_one_logger") +@patch(f"{MOCK_BASE}.get_num_microbatches", return_value=1) +@patch(f"{MOCK_BASE}.print_rank_last") +@patch(f"{MOCK_BASE}.is_megatron_version_bigger_than", return_value=False) +@patch(f"{MOCK_BASE}.torch.cuda.memory_stats", return_value={}) +@patch(f"{MOCK_BASE}.num_floating_point_operations", return_value=1e12) +@patch(f"{MOCK_BASE}.report_memory") +@patch(f"{MOCK_BASE}.report_theoretical_memory") +@patch(f"{MOCK_BASE}.mpu") +@patch(f"{MOCK_BASE}.datetime") +def test_timers_log_when_tensorboard_is_disabled( + mock_datetime, + mock_mpu, + mock_report_theoretical_memory, + mock_report_memory, + mock_num_fp_ops, + mock_mem_stats, + mock_version_check, + mock_print_rank_last, + mock_get_num_microbatches, + mock_get_one_logger, + mock_get_wandb_writer, + mock_get_tensorboard_writer, + mock_get_timers, + mock_get_args, +): + """Case 1: Test timers.log() is called when log_timers_to_tensorboard is False.""" + mock_args, mock_timers_instance, common_kwargs, iteration = _setup_training_log_mocks( + mock_datetime, + mock_mpu, + mock_report_theoretical_memory, + mock_report_memory, + mock_num_fp_ops, + mock_mem_stats, + mock_version_check, + mock_print_rank_last, + mock_get_num_microbatches, + mock_get_one_logger, + mock_get_wandb_writer, + mock_get_tensorboard_writer, + mock_get_timers, + mock_get_args, + ) + + mock_args.log_timers_to_tensorboard = False + training_log(iteration=iteration, **common_kwargs) + mock_timers_instance.log.assert_called_once() + + +@pytest.mark.skipif(not is_megatron_lm_available(), reason="Megatron-LM not available.") +@patch(f"{MOCK_BASE}.get_args") +@patch(f"{MOCK_BASE}.get_timers") +@patch(f"{MOCK_BASE}.get_tensorboard_writer") +@patch(f"{MOCK_BASE}.get_wandb_writer") +@patch(f"{MOCK_BASE}.get_one_logger") +@patch(f"{MOCK_BASE}.get_num_microbatches", return_value=1) +@patch(f"{MOCK_BASE}.print_rank_last") +@patch(f"{MOCK_BASE}.is_megatron_version_bigger_than", return_value=False) +@patch(f"{MOCK_BASE}.torch.cuda.memory_stats", return_value={}) +@patch(f"{MOCK_BASE}.num_floating_point_operations", return_value=1e12) +@patch(f"{MOCK_BASE}.report_memory") +@patch(f"{MOCK_BASE}.report_theoretical_memory") +@patch(f"{MOCK_BASE}.mpu") +@patch(f"{MOCK_BASE}.datetime") +def test_timers_log_when_not_tensorboard_iteration( + mock_datetime, + mock_mpu, + mock_report_theoretical_memory, + mock_report_memory, + mock_num_fp_ops, + mock_mem_stats, + mock_version_check, + mock_print_rank_last, + mock_get_num_microbatches, + mock_get_one_logger, + mock_get_wandb_writer, + mock_get_tensorboard_writer, + mock_get_timers, + mock_get_args, +): + """Case 2: Test timers.log() is called when it's not a tensorboard log iteration.""" + mock_args, mock_timers_instance, common_kwargs, iteration = _setup_training_log_mocks( + mock_datetime, + mock_mpu, + mock_report_theoretical_memory, + mock_report_memory, + mock_num_fp_ops, + mock_mem_stats, + mock_version_check, + mock_print_rank_last, + mock_get_num_microbatches, + mock_get_one_logger, + mock_get_wandb_writer, + mock_get_tensorboard_writer, + mock_get_timers, + mock_get_args, + ) + + mock_args.log_timers_to_tensorboard = True + mock_args.tensorboard_log_interval = 3 # 10 % 3 != 0 + training_log(iteration=iteration, **common_kwargs) + mock_timers_instance.log.assert_called_once() + + +@pytest.mark.skipif(not is_megatron_lm_available(), reason="Megatron-LM not available.") +@patch(f"{MOCK_BASE}.get_args") +@patch(f"{MOCK_BASE}.get_timers") +@patch(f"{MOCK_BASE}.get_tensorboard_writer") +@patch(f"{MOCK_BASE}.get_wandb_writer") +@patch(f"{MOCK_BASE}.get_one_logger") +@patch(f"{MOCK_BASE}.get_num_microbatches", return_value=1) +@patch(f"{MOCK_BASE}.print_rank_last") +@patch(f"{MOCK_BASE}.is_megatron_version_bigger_than", return_value=False) +@patch(f"{MOCK_BASE}.torch.cuda.memory_stats", return_value={}) +@patch(f"{MOCK_BASE}.num_floating_point_operations", return_value=1e12) +@patch(f"{MOCK_BASE}.report_memory") +@patch(f"{MOCK_BASE}.report_theoretical_memory") +@patch(f"{MOCK_BASE}.mpu") +@patch(f"{MOCK_BASE}.datetime") +def test_timers_no_log_on_tensorboard_iteration( + mock_datetime, + mock_mpu, + mock_report_theoretical_memory, + mock_report_memory, + mock_num_fp_ops, + mock_mem_stats, + mock_version_check, + mock_print_rank_last, + mock_get_num_microbatches, + mock_get_one_logger, + mock_get_wandb_writer, + mock_get_tensorboard_writer, + mock_get_timers, + mock_get_args, +): + """Case 3: Test timers.log() is NOT called on a tensorboard log iteration.""" + mock_args, mock_timers_instance, common_kwargs, iteration = _setup_training_log_mocks( + mock_datetime, + mock_mpu, + mock_report_theoretical_memory, + mock_report_memory, + mock_num_fp_ops, + mock_mem_stats, + mock_version_check, + mock_print_rank_last, + mock_get_num_microbatches, + mock_get_one_logger, + mock_get_wandb_writer, + mock_get_tensorboard_writer, + mock_get_timers, + mock_get_args, + ) + + mock_args.log_timers_to_tensorboard = True + mock_args.tensorboard_log_interval = 5 # 10 % 5 == 0 + training_log(iteration=iteration, **common_kwargs) + mock_timers_instance.log.assert_not_called() + + +@pytest.mark.skipif(not is_megatron_lm_available(), reason="Megatron-LM not available.") +@patch(f"{MOCK_BASE}.get_args") +@patch(f"{MOCK_BASE}.get_timers") +@patch(f"{MOCK_BASE}.get_tensorboard_writer") +@patch(f"{MOCK_BASE}.get_wandb_writer") +@patch(f"{MOCK_BASE}.get_one_logger") +@patch(f"{MOCK_BASE}.get_num_microbatches", return_value=1) +@patch(f"{MOCK_BASE}.print_rank_last") +@patch(f"{MOCK_BASE}.is_megatron_version_bigger_than", return_value=False) +@patch(f"{MOCK_BASE}.torch.cuda.memory_stats", return_value={}) +@patch(f"{MOCK_BASE}.num_floating_point_operations", return_value=1e12) +@patch(f"{MOCK_BASE}.report_memory") +@patch(f"{MOCK_BASE}.report_theoretical_memory") +@patch(f"{MOCK_BASE}.mpu") +@patch(f"{MOCK_BASE}.datetime") +@patch("atorch.local_sgd.megatron.parallel_state.get_non_data_parallel_group") +@patch.object(megatron_utils, "is_float8tensor", create=True) +@patch.object(megatron_legacy_model, "Float16Module", create=True) +def test_timers_log_for_local_sgd( + mock_float16module, + mock_is_float8tensor, + mock_get_group, + mock_datetime, + mock_mpu, + mock_report_theoretical_memory, + mock_report_memory, + mock_num_fp_ops, + mock_mem_stats, + mock_version_check, + mock_print_rank_last, + mock_get_num_microbatches, + mock_get_one_logger, + mock_get_wandb_writer, + mock_get_tensorboard_writer, + mock_get_timers, + mock_get_args, +): + """Case 4: Test timers.log() is called correctly for local SGD.""" + mock_args, mock_timers_instance, common_kwargs, iteration = _setup_training_log_mocks( + mock_datetime, + mock_mpu, + mock_report_theoretical_memory, + mock_report_memory, + mock_num_fp_ops, + mock_mem_stats, + mock_version_check, + mock_print_rank_last, + mock_get_num_microbatches, + mock_get_one_logger, + mock_get_wandb_writer, + mock_get_tensorboard_writer, + mock_get_timers, + mock_get_args, + ) + + mock_get_group.return_value = "dummy_group" + mock_args.use_local_sgd = True + mock_args.log_timers_to_tensorboard = False + training_log(iteration=iteration, **common_kwargs) + mock_timers_instance.log.assert_called_once_with( + [ + "forward-backward", + "forward-compute", + "backward-compute", + "batch-generator", + "forward-recv", + "forward-send", + "backward-recv", + "backward-send", + "forward-send-forward-recv", + "forward-send-backward-recv", + "backward-send-forward-recv", + "backward-send-backward-recv", + "forward-backward-send-forward-backward-recv", + "layernorm-grads-all-reduce", + "embedding-grads-all-reduce", + "all-grads-sync", + "params-all-gather", + "optimizer-copy-to-main-grad", + "optimizer-unscale-and-check-inf", + "optimizer-clip-main-grad", + "optimizer-count-zeros", + "optimizer-inner-step", + "optimizer-copy-main-to-model-params", + "optimizer", + ], + normalizer=mock_args.log_interval, + process_group="dummy_group", + ) diff --git a/atorch/tests/common_tests/unordered_dataloader_test.py b/atorch/tests/common_tests/unordered_dataloader_test.py index e1a7185..bd1367c 100644 --- a/atorch/tests/common_tests/unordered_dataloader_test.py +++ b/atorch/tests/common_tests/unordered_dataloader_test.py @@ -1,11 +1,14 @@ import time import unittest +import pytest import torch from torch.utils.data import Dataset, IterableDataset from atorch.data import UnorderedDataLoader +pytest.skip("skips as this test may fail.", allow_module_level=True) + class _TestDataset(Dataset): def __init__(self, data_size=32, sleep_time=0): diff --git a/atorch/tests/common_tests/virtual_optimizer_test.py b/atorch/tests/common_tests/virtual_optimizer_test.py new file mode 100644 index 0000000..85dd384 --- /dev/null +++ b/atorch/tests/common_tests/virtual_optimizer_test.py @@ -0,0 +1,335 @@ +import types +import unittest +from unittest.mock import MagicMock + +import torch + +from atorch.utils.virtual_optimizer.patch_utils import ( + patch_chained_optimizer, + patch_distributed_optimizer, + virtual_distributed_optimizer_load_state_dict, + zero_out_shard_fp32_memory, +) +from atorch.utils.virtual_optimizer.pp_calc import is_valid_pipeline_parallel_combination + + +class VirtualOptimizerTest(unittest.TestCase): + def test_zero_out_shard_fp32_memory(self): + chained_optimizer = MagicMock() + chained_optimizer.chained_optimizers = [MagicMock()] + chained_optimizer.chained_optimizers[0].shard_fp32_groups = [[torch.randn(4)]] + chained_optimizer.chained_optimizers[0].shard_fp32_from_float16_groups = [[torch.randn(4)]] + zero_out_shard_fp32_memory(chained_optimizer) + + self.assertEqual(chained_optimizer.chained_optimizers[0].shard_fp32_groups[0][0].data.shape, (1,)) + self.assertEqual(chained_optimizer.chained_optimizers[0].shard_fp32_from_float16_groups[0][0].data.shape, (1,)) + + def test_patch_chained_optimizer(self): + chained_optimizer = MagicMock() + patch_chained_optimizer(chained_optimizer) + self.assertTrue(hasattr(chained_optimizer, "step")) + self.assertTrue(hasattr(chained_optimizer, "load_state_dict")) + self.assertTrue(hasattr(chained_optimizer, "sharded_state_dict")) + self.assertTrue(hasattr(chained_optimizer, "reload_model_params")) + self.assertTrue(hasattr(chained_optimizer, "load_parameter_state")) + + step_result = chained_optimizer.step() + self.assertEqual(step_result, (True, 0.0, 0)) + + virtual_sharded_state_dict = chained_optimizer.sharded_state_dict(MagicMock()) + self.assertEqual(virtual_sharded_state_dict, {}) + + chained_optimizer.reload_model_params() + chained_optimizer.load_parameter_state(MagicMock()) + + def test_distributed_optimizer_patch_copy_model_grads_to_main_grads(self): + # Mock objects setup + distributed_optimizer = MagicMock() + distributed_optimizer.config = MagicMock() + distributed_optimizer.config.use_precision_aware_optimizer = False + distributed_optimizer.is_stub_optimizer = False + + distributed_optimizer.model_float16_groups = [[MagicMock()]] + distributed_optimizer.shard_fp32_from_float16_groups = [[MagicMock()]] + + patch_distributed_optimizer(distributed_optimizer) + self.assertTrue(hasattr(distributed_optimizer, "_copy_model_grads_to_main_grads")) + + param_range = MagicMock() + param_range.start = 0 + param_range.end = 2 + distributed_optimizer._get_model_param_range_map = MagicMock(return_value={"param": param_range}) + + model_param = torch.randn(4) + model_param.main_grad = torch.randn(4) + distributed_optimizer.model_float16_groups[0][0] = model_param + + shard_main_param = torch.randn(4) + distributed_optimizer.shard_fp32_from_float16_groups[0][0] = shard_main_param + + distributed_optimizer._copy_model_grads_to_main_grads() + self.assertEqual( + distributed_optimizer.shard_fp32_from_float16_groups[0][0].virtual_grad.shape, + model_param.main_grad.view(-1)[param_range.start : param_range.end].shape, + ) + + def test_distributed_optimizer_patch_get_main_grads_for_grad_norm(self): + # Mock objects setup + distributed_optimizer = MagicMock() + distributed_optimizer.config = MagicMock() + distributed_optimizer.config.use_precision_aware_optimizer = False + + patch_distributed_optimizer(distributed_optimizer) + + mock_params = [] + for _ in range(2): # 创建两个参数作为示例 + param = torch.nn.Parameter(torch.randn(2, 2)) + # 添加 virtual_grad 属性 + param.virtual_grad = torch.randn(2, 2) + # 添加 param_is_not_shared 需要的属性 + param.shared = False # 不是共享参数 + # 添加 param_is_not_tensor_parallel_duplicate 需要的属性 + param.tensor_model_parallel = True # 不是重复参数 + mock_params.append(param) + + distributed_optimizer.get_parameters = MagicMock(return_value=mock_params) + self.assertTrue(hasattr(distributed_optimizer, "get_main_grads_for_grad_norm")) + try: + grads = distributed_optimizer.get_main_grads_for_grad_norm() + self.assertIsInstance(grads, list) + self.assertEqual(len(grads), 2) + except Exception as e: + print(f"megatron env may not prepared. error: {e}") + + def test_virtual_distributed_optimizer_load_state_dict(self): + distributed_optimizer = MagicMock() + + distributed_optimizer.load_state_dict = types.MethodType( + virtual_distributed_optimizer_load_state_dict, distributed_optimizer + ) + distributed_optimizer.load_state_dict(MagicMock()) + + self.assertTrue(hasattr(distributed_optimizer, "step")) + + +class UtilsTest(unittest.TestCase): + def test_is_valid_pipeline_parallel_combination(self): + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=12, + pipeline_model_parallel_size=2, + decoder_first_pipeline_num_layers=4, + decoder_last_pipeline_num_layers=4, + decoder_first_virtual_pipeline_num_layers=2, + decoder_last_virtual_pipeline_num_layers=2, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=12, + pipeline_model_parallel_size=3, + decoder_first_pipeline_num_layers=4, + decoder_last_pipeline_num_layers=4, + decoder_first_virtual_pipeline_num_layers=2, + decoder_last_virtual_pipeline_num_layers=2, + num_virtual_stages_per_pipeline_rank=1, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=12, + pipeline_model_parallel_size=3, + decoder_first_pipeline_num_layers=2, + decoder_last_pipeline_num_layers=4, + decoder_first_virtual_pipeline_num_layers=1, + decoder_last_virtual_pipeline_num_layers=2, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=12, + pipeline_model_parallel_size=3, + decoder_first_pipeline_num_layers=4, + decoder_last_pipeline_num_layers=2, + decoder_first_virtual_pipeline_num_layers=2, + decoder_last_virtual_pipeline_num_layers=1, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=13, + pipeline_model_parallel_size=4, + decoder_first_pipeline_num_layers=4, + decoder_last_pipeline_num_layers=4, + decoder_first_virtual_pipeline_num_layers=2, + decoder_last_virtual_pipeline_num_layers=2, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=18, + pipeline_model_parallel_size=6, + decoder_first_pipeline_num_layers=6, + decoder_last_pipeline_num_layers=3, + decoder_first_virtual_pipeline_num_layers=3, + decoder_last_virtual_pipeline_num_layers=1, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=18, + pipeline_model_parallel_size=6, + decoder_first_pipeline_num_layers=3, + decoder_last_pipeline_num_layers=6, + decoder_first_virtual_pipeline_num_layers=1, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=64, + pipeline_model_parallel_size=3, + decoder_first_pipeline_num_layers=6, + decoder_last_pipeline_num_layers=6, + decoder_first_virtual_pipeline_num_layers=3, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=24, + pipeline_model_parallel_size=4, + decoder_first_pipeline_num_layers=6, + decoder_last_pipeline_num_layers=6, + decoder_first_virtual_pipeline_num_layers=3, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=3, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=24, + pipeline_model_parallel_size=4, + decoder_first_pipeline_num_layers=7, + decoder_last_pipeline_num_layers=6, + decoder_first_virtual_pipeline_num_layers=3, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=3, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=32, + pipeline_model_parallel_size=5, + decoder_first_pipeline_num_layers=10, + decoder_last_pipeline_num_layers=6, + decoder_first_virtual_pipeline_num_layers=2, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=26, + pipeline_model_parallel_size=4, + decoder_first_pipeline_num_layers=8, + decoder_last_pipeline_num_layers=6, + decoder_first_virtual_pipeline_num_layers=5, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=24, + pipeline_model_parallel_size=4, + decoder_first_pipeline_num_layers=8, + decoder_last_pipeline_num_layers=6, + decoder_first_virtual_pipeline_num_layers=5, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=32, + pipeline_model_parallel_size=4, + decoder_first_pipeline_num_layers=8, + decoder_last_pipeline_num_layers=6, + decoder_first_virtual_pipeline_num_layers=1, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=24, + pipeline_model_parallel_size=4, + decoder_first_pipeline_num_layers=6, + decoder_last_pipeline_num_layers=7, + decoder_first_virtual_pipeline_num_layers=3, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=3, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=30, + pipeline_model_parallel_size=5, + decoder_first_pipeline_num_layers=6, + decoder_last_pipeline_num_layers=10, + decoder_first_virtual_pipeline_num_layers=3, + decoder_last_virtual_pipeline_num_layers=2, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertFalse( + is_valid_pipeline_parallel_combination( + num_layers=32, + pipeline_model_parallel_size=4, + decoder_first_pipeline_num_layers=6, + decoder_last_pipeline_num_layers=8, + decoder_first_virtual_pipeline_num_layers=3, + decoder_last_virtual_pipeline_num_layers=6, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + self.assertTrue( + is_valid_pipeline_parallel_combination( + num_layers=24, + pipeline_model_parallel_size=4, + decoder_first_pipeline_num_layers=6, + decoder_last_pipeline_num_layers=6, + decoder_first_virtual_pipeline_num_layers=3, + decoder_last_virtual_pipeline_num_layers=3, + num_virtual_stages_per_pipeline_rank=2, + ) + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/atorch/tests/toy_modules/toy_module_te.py b/atorch/tests/toy_modules/toy_module_te.py new file mode 100644 index 0000000..887b455 --- /dev/null +++ b/atorch/tests/toy_modules/toy_module_te.py @@ -0,0 +1,61 @@ +import math + +import torch +import torch.nn as nn + +try: + from transformer_engine.pytorch import GroupedLinear, Linear + + HAS_TE = True +except (ImportError, ModuleNotFoundError): + GroupedLinear = object + Linear = object + HAS_TE = False + + +# Toy for moe +class DummyTeGG(torch.nn.Module): + def __init__(self, hidden_size, intermediate_size, num_gemms): + super().__init__() + self.gg1 = GroupedLinear(num_gemms, hidden_size, intermediate_size) + self.gg2 = GroupedLinear(num_gemms, intermediate_size, hidden_size) + self.num_gemms = num_gemms + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + + def forward(self, x): + token_num = math.prod(x.shape[:-1]) + # hack to evenly partition input x + m_splits = [token_num // self.num_gemms] * self.num_gemms + return self.gg2(self.gg1(x, m_splits), m_splits) + + +class DummyTeModel(torch.nn.Module): + def __init__(self, hidden_size, num_gemms=4): + super().__init__() + self.mw1 = Linear(hidden_size, hidden_size * 2) + self.mw2 = Linear(hidden_size * 2, hidden_size) + self.layers = torch.nn.ModuleList(DummyTeGG(hidden_size, 2 * hidden_size, num_gemms) for _ in range(2)) + + def forward(self, x): + x = self.mw2(self.mw1(x)) + for layer in self.layers: + x = layer(x) + return x + + +def get_model(hidden_size, num_gemms=4, seed=123): + torch.cuda.manual_seed(seed) + torch.manual_seed(seed) + + model = DummyTeModel(hidden_size, num_gemms).cuda() + return model + + +def get_input(batch_size, hidden_size): + return torch.randn(batch_size, hidden_size, device=torch.device("cuda")) + + +def loss_func(inputs, output): + loss = nn.MSELoss() + return loss(inputs, output) diff --git a/atorch/tests/utils/test_tools.py b/atorch/tests/utils/test_tools.py new file mode 100644 index 0000000..b78c2a9 --- /dev/null +++ b/atorch/tests/utils/test_tools.py @@ -0,0 +1,64 @@ +import tempfile +import unittest + +import torch + +from atorch.utils.parse_memory_pickle import parse_args as parse_memory_pickle_args +from atorch.utils.parse_memory_pickle import parse_memory_pickle_file, print_result +from atorch.utils.parse_trace_json import parse_trace_file, print_profiler_summary + + +class TestModule(torch.nn.Module): + def __init__(self): + super().__init__() + self.linear = torch.nn.Linear(10, 10) + + def forward(self, x): + return self.linear(x) + + +@unittest.skipIf(not torch.cuda.is_available(), "CUDA is not available") +class TestTools(unittest.TestCase): + def setUp(self): + self.store_dir = tempfile.TemporaryDirectory() + torch.cuda.memory._record_memory_history() + + def tearDown(self): + self.store_dir.cleanup() + torch.cuda.memory._record_memory_history(enabled=None) + + def gen_trace_file(self): + linear = TestModule() + linear.cuda() + x = torch.randn(10, 10).cuda() + + with torch.profiler.profile( + activities=[ + torch.profiler.ProfilerActivity.CPU, + torch.profiler.ProfilerActivity.CUDA, + ], + with_stack=True, # python stack is empty? + record_shapes=True, + ) as prof: + out = [] + for _ in range(100): + y = linear(x) + out.append(y) + profiler_file = f"{self.store_dir.name}/trace.json" + prof.export_chrome_trace(profiler_file) + memory_pickle_file = f"{self.store_dir.name}/memory.pickle" + torch.cuda.memory._dump_snapshot(memory_pickle_file) + return profiler_file, memory_pickle_file + + def test_parse_trace_file(self): + profiler_file, memory_pickle_file = self.gen_trace_file() + summary, kernel_start_time = parse_trace_file(profiler_file) + print_profiler_summary([summary], [kernel_start_time]) + + def test_parse_memory_pickle_file(self): + profiler_file, memory_pickle_file = self.gen_trace_file() + args = parse_memory_pickle_args([memory_pickle_file, "--whitelist", "atorch", "torch"]) + blocklist = args.blocklist or [] + whitelist = args.whitelist or [] + filename_counter, final_counter = parse_memory_pickle_file(memory_pickle_file, whitelist) + print_result(filename_counter, blocklist, whitelist) diff --git a/atorch/trainer/args.py b/atorch/trainer/args.py index 237e0f3..8163d78 100644 --- a/atorch/trainer/args.py +++ b/atorch/trainer/args.py @@ -27,9 +27,9 @@ "eval_steps": "eval_interval", "test_steps": "test_interval", "tensorboard_dir": "tensorboard_dir", - # "manual_gc": "manual_gc", - # "manual_gc_interval": "manual_gc_interval", - # "manual_gc_eval": "manual_gc_eval", + "manual_gc": "manual_gc", + "manual_gc_interval": "manual_gc_interval", + "manual_gc_eval": "manual_gc_eval", # torch init process # Comment the following two lines temporarily. # "ddp_backend": "distributed_backend", @@ -79,10 +79,15 @@ class AtorchTrainingArgs(DataclassMixin, AutoMapperExtraConfigs): debug_switch: dict = field( default_factory=dict, metadata={ - "help": "switches to enable some features for debugging, " "supported features: " " - log_torch_save: True" + "help": "switches to enable some features for debugging, " + "supported features: " + " - log_torch_save: True " + " - print_config: True" }, ) + debug_module: bool = field(default=False, metadata={"help": "Open module debug tool."}) + flash_checkpoint: bool = field( default=False, metadata={ @@ -134,6 +139,7 @@ class AtorchTrainingArgs(DataclassMixin, AutoMapperExtraConfigs): default=None, metadata={"help": "The path to a folder with a valid checkpoint for your model."}, ) + resume_strict: Optional[bool] = field(default=True, metadata={"help": "Whether to load checkpoint strictly."}) do_train: bool = field(default=False, metadata={"help": "Whether to run training."}) do_eval: bool = field(default=False, metadata={"help": "Whether to run eval on the dev set."}) @@ -257,6 +263,8 @@ class AtorchTrainingArgs(DataclassMixin, AutoMapperExtraConfigs): "Virtualize optimizer to reduce the number of GPUs required for validating hybrid parallel strategies." }, ) + memory_snapshot_path: Optional[str] = field(default=None, metadata={"help": "The path to save memory snapshot."}) + memory_snapshot_step: Optional[int] = field(default=1, metadata={"help": "The step to save memory snapshot."}) safe_serialization: Optional[bool] = field(default=True) @@ -285,6 +293,15 @@ class AtorchTrainingArgs(DataclassMixin, AutoMapperExtraConfigs): default=None, metadata={"help": "Save frequency in each epoch."} ) + """ file content: + { + "save_at_dynamic_steps": [1000] + } + """ + dynamic_save_config_path: Optional[str] = field( + default=None, metadata={"help": "The file path to monitor whether to save checkpoint at any time."} + ) + # Evaluation evaluation_strategy: Union[IntervalStrategy, str] = field( default="no", @@ -589,7 +606,7 @@ def __post_init__(self): raise ValueError("`gradient_accumulation_steps` is invalid in Megatron training mode.") if self.finetune_type == "dpo" and self.custom_dpo_infer_function is None: - raise ValueError("Arg `custom_dpo_infer_function` should be provided when training DPO task.") + logger.warning("Arg `custom_dpo_infer_function` should be provided when training DPO task.") if len(self.debug_switch) > 0: log_torch_save = self.debug_switch.get("log_torch_save", False) @@ -841,6 +858,20 @@ class MegatronArgs(DynamicDataClass): metadata={"help": "Custom args, as a complement for Megatron-LM arguments."}, ) + # args for function megatron/training/initialize::initialize_megatron() + get_embedding_ranks: Optional[Callable] = field( + default=None, + metadata={"help": "A function to get the embedding ranks."}, + ) + get_position_embedding_ranks: Optional[Callable] = field( + default=None, + metadata={"help": "A function to get the position embedding ranks."}, + ) + get_output_layer_ranks: Optional[Callable] = field( + default=None, + metadata={"help": "A function to get the output layer ranks."}, + ) + def __init__(self, **kwargs): super().__init__(**kwargs) diff --git a/atorch/trainer/atorch_trainer_v2.py b/atorch/trainer/atorch_trainer_v2.py index 60deaeb..837b875 100644 --- a/atorch/trainer/atorch_trainer_v2.py +++ b/atorch/trainer/atorch_trainer_v2.py @@ -5,7 +5,6 @@ import random import sys import time -from contextlib import nullcontext from pathlib import Path from typing import Dict, List, Optional, Tuple, Union @@ -32,7 +31,6 @@ from atorch.common.log_utils import default_logger as logger from atorch.distributed.distributed import is_distributed from atorch.trainer.args import AtorchTrainingArgs -from atorch.trainer.atorch_profiler import get_profiler from atorch.trainer.base.atorch_container import AtorchTrainerContainer from atorch.trainer.base.atorch_module import AtorchIRModel from atorch.trainer.base.atorch_train_engine import AtorchTrainEngine @@ -44,15 +42,22 @@ AtorchTrainerState, FlowCallbackV2, PredictCallback, + ProfilerCallback, ) from atorch.trainer.utils import DistributedType, print_all_args_before_training from atorch.utils.hooks import ATorchHooks from atorch.utils.import_util import is_megatron_lm_available, is_torch_npu_available if is_megatron_lm_available(): + from megatron.core.num_microbatches_calculator import get_current_global_batch_size from megatron.training.global_vars import get_args, get_timers - from atorch.trainer.megatron import AtorchMegatronEngine + from atorch.trainer.debug_utils.debug_module import DebugCallback + from atorch.trainer.megatron import ( + AtorchMegatronEngine, + MegatronCallback, + skip_first_batches_for_megatron_dataloader, + ) if is_torch_npu_available(): try: @@ -60,7 +65,6 @@ except ImportError: dp = None - DEFAULT_FLOW_CALLBACKS = [FlowCallbackV2] DEFAULT_PROGRESS_CALLBACK = ProgressCallback @@ -97,6 +101,13 @@ def recursive_jsons_to_dict(value): return value +def _record_memory_snapshot(args): + if not os.path.exists(args.memory_snapshot_path): + os.makedirs(args.memory_snapshot_path, exist_ok=True) + torch.cuda.memory._dump_snapshot(f"{args.memory_snapshot_path}/snap_{torch.distributed.get_rank()}.pickle") + torch.cuda.memory._record_memory_history(enabled=None) + + class AtorchTrainerV2: def __init__( self, @@ -158,6 +169,9 @@ def __init__( self.callback_handler = AtorchCallbackHandler( callbacks, self.model, self.tokenizer, self.optimizer, self.lr_scheduler ) + + self.add_callback(ProfilerCallback) + if self.distributed_type == DistributedType.MEGATRON: self.add_callback(PredictCallback) else: @@ -201,10 +215,6 @@ def pop_callback(self, callback: TrainerCallback): """ return self.callback_handler.pop_callback(callback) - @property - def use_distributed(self): - return self.distributed_type != DistributedType.NO_DISTRIBUTE - def train( self, **kwargs, @@ -226,10 +236,11 @@ def train( ) resume_from_checkpoint = self.args.resume_from_checkpoint + args = self.args - return self._inner_training_loop(args=self.args, resume_from_checkpoint=resume_from_checkpoint) + if self.args.memory_snapshot_path is not None: + torch.cuda.memory._record_memory_history(max_entries=10000000) - def _inner_training_loop(self, args: AtorchTrainingArgs, resume_from_checkpoint=None, **kwargs): if self.datasets is not None: dataset_num = len(self.datasets) @@ -237,10 +248,6 @@ def _inner_training_loop(self, args: AtorchTrainingArgs, resume_from_checkpoint= eval_dataset = self.datasets[1] if dataset_num > 1 else None test_dataset = self.datasets[2] if dataset_num > 2 else None - # Log a few random samples from the training set: - # for index in random.sample(range(len(train_dataset)), 3): - # logger.info(f"Sample {index} of the training set: {train_dataset[index]}.") - train_dataloader = None eval_dataloader = None test_dataloader = None @@ -284,104 +291,39 @@ def _inner_training_loop(self, args: AtorchTrainingArgs, resume_from_checkpoint= elif self.distributed_type == DistributedType.MEGATRON: if not is_megatron_lm_available(): raise ValueError("Megatron-LM is not installed.") - # TODO: check type self.train_engine = AtorchMegatronEngine( - # TODO YY: adjust input train_args=args, - model=self.model, - optimizer=self.optimizer, - scheduler=self.lr_scheduler, train_state=self.state, dataloaders=dataloaders, resume_from_checkpoint=resume_from_checkpoint, **kwargs, ) megatron_args = get_args() - else: - raise NotImplementedError(f"Not implemented distributed backend {self.distributed_type}.") - try: - from ant_utils.dcp_utils.dcp_utils import patch_torch + # add MegatronCallback + self.add_callback(MegatronCallback(self.train_engine)) - patch_torch(patch_find_nd_overlapping_shards=self.args.patch_find_nd_overlapping_shards) - except ImportError: - logger.warning( - "Unable to import patch_torch from ant_utils.dcp_utils.dcp_utils, you might not using " - "available megatron version. If you want to use a speedup version of megatron," - "please read atorch doc to use a proper megatron version." - ) + if self.args.debug_module: + self.add_callback(DebugCallback(self.train_engine)) + + else: + raise NotImplementedError(f"Not implemented distributed backend {self.distributed_type}.") train_dataloader = self.train_engine.get_dataloader("train") - if self.args.finetune_type == "dpo": - self.args.custom_dpo_infer_function(self.train_engine.module, self.train_engine._dataloaders) + if self.args.finetune_type == "dpo" and self.args.custom_dpo_infer_function is not None: + self.args.custom_dpo_infer_function(self.train_engine) total_train_batch_size = args.global_train_batch_size * args.gradient_accumulation_steps - len_dataloader = None - if self.distributed_type == DistributedType.MEGATRON: - if self.args.finetune_type is not None: - len_dataloader = len(train_dataloader) - num_update_steps_per_epoch = max(len_dataloader, 1) - num_examples = train_dataloader.num_examples - - # Internal checking - # args.max_steps has been calculated in MegatronEngine - assert args.max_steps > 0, "`max_steps` should be greater than 0" - assert ( - args.max_steps == megatron_args.train_iters - ), "Please ensure trainer args.max_steps is equal to megatron_args.train_iters." - max_steps = args.max_steps - num_train_epochs = args.max_steps // num_update_steps_per_epoch + int( - args.max_steps % num_update_steps_per_epoch > 0 - ) - # May be slightly incorrect if the last batch in the training dataloader has a smaller size but it's - # the best we can do. - num_train_samples = max_steps * total_train_batch_size - else: - if ( - args.max_steps > 0 - and megatron_args.train_iters is not None - and args.max_steps != megatron_args.train_iters - ): - logger.warning( - "args.max_steps will be overwritten by megatron_args.train_iters under MEGATRON training mode." - ) - args.max_steps = megatron_args.train_iters - max_steps = args.max_steps - num_train_epochs = sys.maxsize - num_update_steps_per_epoch = max_steps - num_examples = total_train_batch_size * max_steps - num_train_samples = max_steps * total_train_batch_size - elif has_length(train_dataloader): - len_dataloader = len(train_dataloader) - num_update_steps_per_epoch = len_dataloader // args.gradient_accumulation_steps - num_update_steps_per_epoch = max(num_update_steps_per_epoch, 1) - num_examples = self.num_examples(train_dataloader) - if args.max_steps > 0: - max_steps = args.max_steps - num_train_epochs = args.max_steps // num_update_steps_per_epoch + int( - args.max_steps % num_update_steps_per_epoch > 0 - ) - # May be slightly incorrect if the last batch in the training dataloader has a smaller size but it's - # the best we can do. - num_train_samples = args.max_steps * total_train_batch_size - else: - max_steps = math.ceil(args.num_train_epochs * num_update_steps_per_epoch) - num_train_epochs = math.ceil(args.num_train_epochs) - num_train_samples = self.num_examples(train_dataloader) * args.num_train_epochs # type: ignore[assignment] # noqa: E501 - elif args.max_steps > 0: # Rely on max_steps when dataloader does not have a working size - max_steps = args.max_steps - # Setting a very large number of epochs so we go as many times as necessary over the iterator. - num_train_epochs = sys.maxsize - num_update_steps_per_epoch = max_steps - num_examples = total_train_batch_size * args.max_steps - num_train_samples = args.max_steps * total_train_batch_size - else: - raise ValueError( - "args.max_steps must be set to a positive value if dataloader does not have a length, was" - f" {args.max_steps}" - ) + ( + len_dataloader, + max_steps, + num_train_epochs, + num_update_steps_per_epoch, + num_examples, + num_train_samples, + ) = self._get_train_num(args, train_dataloader, total_train_batch_size) # Compute absolute values for logging, eval, and save if given as ratio if args.logging_steps and args.logging_steps < 1: @@ -391,13 +333,6 @@ def _inner_training_loop(self, args: AtorchTrainingArgs, resume_from_checkpoint= if args.save_steps and args.save_steps < 1: args.save_steps = math.ceil(max_steps * args.save_steps) - # Not support search hyper param - # self.state.is_hyper_param_search = False - # if package_version_bigger_than("transformers", "4.31.0"): - # self.state.logging_steps = args.logging_steps - # self.state.eval_steps = args.eval_steps - # self.state.save_steps = args.save_steps - # Train! logger.info("***** Running training *****") logger.info(f" Num examples = {num_examples:,}") @@ -440,10 +375,17 @@ def _inner_training_loop(self, args: AtorchTrainingArgs, resume_from_checkpoint= ) # Update the references - self.callback_handler.model = self.model # TODO: Update to engine - self.callback_handler.optimizer = self.optimizer - self.callback_handler.lr_scheduler = self.lr_scheduler - self.callback_handler.train_dataloader = train_dataloader + if self.distributed_type == DistributedType.MEGATRON: + self.callback_handler.model = self.train_engine.module + self.callback_handler.optimizer = self.train_engine.optimizer + self.callback_handler.lr_scheduler = self.train_engine.scheduler + self.callback_handler.train_dataloader = train_dataloader + self.callback_handler.eval_dataloader = self.train_engine.get_dataloader("eval") + else: + self.callback_handler.model = self.model + self.callback_handler.optimizer = self.optimizer + self.callback_handler.lr_scheduler = self.lr_scheduler + self.callback_handler.train_dataloader = train_dataloader # This should be the same if the state has been saved but in case the training arguments changed, it's safer # to set this after the load. self.state.max_steps = max_steps @@ -483,6 +425,8 @@ def _inner_training_loop(self, args: AtorchTrainingArgs, resume_from_checkpoint= dist.barrier() + if self.args.finetune_type == "dpo": + self.evaluate_dpo(step=self.state.global_step) # train begin after all ranks reach here self.control = self.callback_handler.on_train_begin(args, self.state, self.control) @@ -493,183 +437,148 @@ def _inner_training_loop(self, args: AtorchTrainingArgs, resume_from_checkpoint= all_args_to_log["MegatronArgs"] = vars(megatron_args) print_all_args_before_training(all_args_to_log) - with get_profiler(self.args) as prof: - total_batched_samples = 0 - for epoch in range(epochs_trained, num_train_epochs): - epoch_iterator = train_dataloader - if self.args.finetune_type is not None: - if hasattr(epoch_iterator, "set_epoch"): - epoch_iterator.set_epoch(epoch) - elif hasattr(epoch_iterator.sampler, "set_epoch"): - epoch_iterator.sampler.set_epoch(epoch) - - steps_in_epoch = ( - len(epoch_iterator) - if len_dataloader is not None - else args.max_steps * args.gradient_accumulation_steps - ) + total_batched_samples = 0 + for epoch in range(epochs_trained, num_train_epochs): + epoch_iterator = train_dataloader + if self.args.finetune_type is not None: + if hasattr(epoch_iterator, "set_epoch"): + epoch_iterator.set_epoch(epoch) + elif hasattr(epoch_iterator.sampler, "set_epoch"): + epoch_iterator.sampler.set_epoch(epoch) - if ( - self.args.extra_save_frequency_in_epoch is not None - and len(self.args.extra_save_frequency_in_epoch) > 0 - ): - if isinstance(self.args.extra_save_frequency_in_epoch[0], float): - self.args.extra_save_frequency_in_epoch = [ - int(f * steps_in_epoch) for f in self.args.extra_save_frequency_in_epoch - ] - logger.info( - f"In this epoch, checkpoint in step {self.args.extra_save_frequency_in_epoch} will be saved! " - f"There are {steps_in_epoch} steps in this epoch." - ) - self.state.steps_in_epoch = steps_in_epoch + steps_in_epoch = ( + len(epoch_iterator) if len_dataloader is not None else args.max_steps * args.gradient_accumulation_steps + ) - self.control = self.callback_handler.on_epoch_begin(args, self.state, self.control) + if self.args.extra_save_frequency_in_epoch is not None and len(self.args.extra_save_frequency_in_epoch) > 0: + if isinstance(self.args.extra_save_frequency_in_epoch[0], float): + self.args.extra_save_frequency_in_epoch = [ + int(f * steps_in_epoch) for f in self.args.extra_save_frequency_in_epoch + ] + logger.info( + f"In this epoch, checkpoint in step {self.args.extra_save_frequency_in_epoch} will be saved! " + f"There are {steps_in_epoch} steps in this epoch." + ) + self.state.steps_in_epoch = steps_in_epoch - if ( - epoch == epochs_trained - and resume_from_checkpoint is not None - and steps_trained_in_current_epoch == 0 - ): - self._load_rng_state(resume_from_checkpoint) + self.control = self.callback_handler.on_epoch_begin(args, self.state, self.control) - rng_to_sync = False - steps_skipped = 0 + if epoch == epochs_trained and resume_from_checkpoint is not None and steps_trained_in_current_epoch == 0: + self._load_rng_state(resume_from_checkpoint) + + rng_to_sync = False + steps_skipped = 0 + if steps_trained_in_current_epoch > 0: if self.distributed_type != DistributedType.MEGATRON: - if steps_trained_in_current_epoch > 0: - epoch_iterator = skip_first_batches(epoch_iterator, steps_trained_in_current_epoch) + epoch_iterator = skip_first_batches(epoch_iterator, steps_trained_in_current_epoch) + steps_skipped = steps_trained_in_current_epoch + steps_trained_in_current_epoch = 0 + rng_to_sync = True + else: + if self.args.finetune_type is None: # Pretrain steps_skipped = steps_trained_in_current_epoch steps_trained_in_current_epoch = 0 - rng_to_sync = True - else: - if steps_trained_in_current_epoch > 0: - # TODO: steps_skipped should be get_args().iteration, because - # self.train_engine.iteration will be real-time updated. - steps_skipped = self.train_engine.iteration - - step = -1 - for step, inputs in enumerate(epoch_iterator): - # TODO: Move it out - self.train_engine.train() - - if ( - self.args.profiler_type == "nsys" - and self.state.global_step == self.args.profile_step_start - and self.args.process_index in self.args.profile_ranks - ): - torch.cuda.cudart().cudaProfilerStart() - torch.autograd.profiler.emit_nvtx(record_shapes=True).__enter__() - - total_batched_samples += 1 - if rng_to_sync: - self._load_rng_state(resume_from_checkpoint) - rng_to_sync = False - - # TODO: remove them. - # # Skip past any already trained steps if resuming training - # if steps_trained_in_current_epoch > 0: - # steps_trained_in_current_epoch -= 1 - # if steps_trained_progress_bar is not None: - # steps_trained_progress_bar.update(1) - # if steps_trained_in_current_epoch == 0: - # self._load_rng_state(resume_from_checkpoint) - # continue - # elif steps_trained_progress_bar is not None: - # steps_trained_progress_bar.close() - # steps_trained_progress_bar = None - - if step % args.gradient_accumulation_steps == 0: - self.control = self.callback_handler.on_step_begin(args, self.state, self.control) - - if self.distributed_type == DistributedType.MEGATRON: - # MegatronEngine's train_step() contains: - # forward, backward, optimizer.step(), zero_grad() - loss = self.train_engine(inputs) - self.state.consumed_train_samples = megatron_args.consumed_train_samples - self.state.consumed_train_tokens = ( - megatron_args.consumed_train_samples * megatron_args.seq_length + else: # Post-training + epoch_iterator = skip_first_batches_for_megatron_dataloader( + epoch_iterator, steps_trained_in_current_epoch ) - self.state.total_flos = self.train_engine.num_floating_point_operations_so_far - else: - with self.train_engine.accumulate(): - loss = self.train_engine(**inputs) - - self.train_engine.backward(loss) - - self.train_engine.optimizer_step() - self.train_engine.scheduler_step() - self.train_engine.optimizer_zero_grad() - self.state.consumed_train_samples += self.args.global_train_batch_size - # TODO: record consumed_train_tokens - - if args.logging_nan_inf_filter and (torch.isnan(loss) or torch.isinf(loss)): - # if loss is nan or inf simply add the average of previous logged losses - tr_loss += tr_loss / (1 + self.state.global_step - self._globalstep_last_logged) - else: - tr_loss += loss - - is_last_step_and_steps_less_than_grad_acc = ( - steps_in_epoch <= args.gradient_accumulation_steps and (step + 1) == steps_in_epoch - ) + steps_skipped = steps_trained_in_current_epoch + steps_trained_in_current_epoch = 0 + # TODO: Moving the operation of sync RNG state from AtorchMegatronEngine here + # rng_to_sync = True - if ( - # last step in epoch but step is always smaller than gradient_accumulation_steps - total_batched_samples % args.gradient_accumulation_steps == 0 - or is_last_step_and_steps_less_than_grad_acc - ): - self.state.global_step += 1 - self.state.epoch = epoch + (step + 1 + steps_skipped) / steps_in_epoch - self.state.current_step_in_epoch = step + 1 + steps_skipped - self.control = self.callback_handler.on_step_end(args, self.state, self.control) - - # TODO: mix log,save and evaluate in one function is not a good practice, split and control - # separately. - self._maybe_log_save_evaluate(tr_loss, self.model, epoch) - else: - self.control = self.callback_handler.on_substep_end(args, self.state, self.control) - - if ( - self.args.profiler_type == "nsys" - and self.state.global_step == self.args.profile_step_end - and self.args.process_index in self.args.profile_ranks - ): - torch.cuda.cudart().cudaProfilerStop() - elif self.args.profiler_type == "hw_dp": - if dp is not None: - dp.step() - elif prof is not None and not isinstance(prof, nullcontext): - prof.step() - - if ( - self.args.empty_cache_steps is not None - and self.state.global_step % self.args.empty_cache_steps == 0 - ): - torch.cuda.empty_cache() - - if self.args.manual_gc: - if ( - self.args.manual_gc_interval != 0 - and self.state.global_step % self.args.manual_gc_interval == 0 - ): - logger.info("Execute GC manually.") - gc.collect() - - if self.control.should_epoch_stop or self.control.should_training_stop: - break - - if step < 0: - logger.warning( - "There seems to be not a single sample in your epoch_iterator, stopping training at step" - f" {self.state.global_step}! This is expected if you're using an IterableDataset and set" - f" num_steps ({max_steps}) higher than the number of available samples." - ) - self.control.should_training_stop = True + step = -1 + for step, inputs in enumerate(epoch_iterator): + # TODO: Move it out + self.train_engine.train() + + total_batched_samples += 1 + if rng_to_sync: + self._load_rng_state(resume_from_checkpoint) + rng_to_sync = False + + if step % args.gradient_accumulation_steps == 0: + self.control = self.callback_handler.on_step_begin(args, self.state, self.control) + + if self.distributed_type == DistributedType.MEGATRON: + # MegatronEngine's train_step() contains: + # forward, backward, optimizer.step(), zero_grad() + + try: + loss = self.train_engine(inputs) + except Exception as e: + if self.args.memory_snapshot_path is not None: + _record_memory_snapshot(self.args) + raise e + if self.args.memory_snapshot_path is not None and step == self.args.memory_snapshot_step: + _record_memory_snapshot(self.args) + self.state.consumed_train_samples = megatron_args.consumed_train_samples + self.state.consumed_train_tokens += get_current_global_batch_size() * megatron_args.seq_length + self.state.total_flos = self.train_engine.num_floating_point_operations_so_far + else: + with self.train_engine.accumulate(): + loss = self.train_engine(**inputs) + + self.train_engine.backward(loss) - self.control = self.callback_handler.on_epoch_end(args, self.state, self.control) - self._maybe_log_save_evaluate(tr_loss, self.model, epoch, epoch_end=True) + self.train_engine.optimizer_step() + self.train_engine.scheduler_step() + self.train_engine.optimizer_zero_grad() + self.state.consumed_train_samples += self.args.global_train_batch_size + # TODO: record consumed_train_tokens - if self.control.should_training_stop: + if args.logging_nan_inf_filter and (torch.isnan(loss) or torch.isinf(loss)): + # if loss is nan or inf simply add the average of previous logged losses + tr_loss += tr_loss / (1 + self.state.global_step - self._globalstep_last_logged) + else: + tr_loss += loss + + is_last_step_and_steps_less_than_grad_acc = ( + steps_in_epoch <= args.gradient_accumulation_steps and (step + 1) == steps_in_epoch + ) + + if ( + # last step in epoch but step is always smaller than gradient_accumulation_steps + total_batched_samples % args.gradient_accumulation_steps == 0 + or is_last_step_and_steps_less_than_grad_acc + ): + self.state.global_step += 1 + self.state.epoch = epoch + (step + 1 + steps_skipped) / steps_in_epoch + self.state.current_step_in_epoch = step + 1 + steps_skipped + self.control = self.callback_handler.on_step_end(args, self.state, self.control) + + self._maybe_log_save_evaluate(tr_loss, self.model, epoch) + else: + self.control = self.callback_handler.on_substep_end(args, self.state, self.control) + + if ( + self.args.empty_cache_steps is not None + and self.state.global_step % self.args.empty_cache_steps == 0 + ): + torch.cuda.empty_cache() + + if self.args.manual_gc: + if self.args.manual_gc_interval != 0 and self.state.global_step % self.args.manual_gc_interval == 0: + logger.info("Execute GC manually.") + gc.collect() + + if self.control.should_epoch_stop or self.control.should_training_stop: break + if step < 0: + logger.warning( + "There seems to be not a single sample in your epoch_iterator, stopping training at step" + f" {self.state.global_step}! This is expected if you're using an IterableDataset and set" + f" num_steps ({max_steps}) higher than the number of available samples." + ) + self.control.should_training_stop = True + + self.control = self.callback_handler.on_epoch_end(args, self.state, self.control) + self._maybe_log_save_evaluate(tr_loss, self.model, epoch, epoch_end=True) + + if self.control.should_training_stop: + break + # add remaining tr_loss self._total_loss_scalar += tr_loss.item() train_loss = self._total_loss_scalar / self.state.global_step @@ -686,13 +595,10 @@ def _inner_training_loop(self, args: AtorchTrainingArgs, resume_from_checkpoint= def _maybe_log_save_evaluate(self, tr_loss, model, epoch, epoch_end=False): if self.distributed_type == DistributedType.MEGATRON: timers = get_timers() - - if self.distributed_type == DistributedType.MEGATRON: logging_metrics = self.train_engine.training_log() if self.control.should_log: self.log(logging_metrics) - - if self.control.should_log and self.distributed_type != DistributedType.MEGATRON: + elif self.control.should_log: logs: Dict[str, float] = {} # all_gather + mean() to get average loss over all processes @@ -715,26 +621,10 @@ def _eval(eval_or_test): # Collect all objects. gc.collect() if self.distributed_type == DistributedType.MEGATRON: - megatron_args = get_args() - timers("interval-time").stop() - if megatron_args.use_distributed_optimizer and megatron_args.overlap_param_gather: - try: - from megatron.training.training import disable_forward_pre_hook # adapt for megatron 0.10 - - disable_forward_pre_hook(self.train_engine.module) - except ImportError: - self.train_engine.optimizer.disable_pre_hook() - - metrics = self.evaluate(eval_or_test=eval_or_test) - if megatron_args.use_distributed_optimizer and megatron_args.overlap_param_gather: - try: - from megatron.training.training import enable_forward_pre_hook # adapt for megatron 0.10 - - enable_forward_pre_hook(self.train_engine.module) - except ImportError: - self.train_engine.optimizer.enable_pre_hook() - - timers("interval-time", log_level=0).start(barrier=True) + if self.args.finetune_type == "dpo": + self.evaluate_dpo(step=self.state.global_step) + else: + metrics = self.evaluate(eval_or_test=eval_or_test) else: # TODO: Implement evaluate metrics = self.evaluate(eval_or_test=eval_or_test) @@ -749,28 +639,30 @@ def _eval(eval_or_test): # Collect only the objects created and used in evaluation. gc.collect(generation=0) + if self.distributed_type == DistributedType.MEGATRON: + timers("interval-time").stop() + if self.control.should_evaluate: _eval(eval_or_test="eval") if self.control.should_test: _eval(eval_or_test="test") - # self.train_engine.post_training_step() - if self.control.should_save: # Save model checkpoint - self.control = self.callback_handler.on_save_begin(self.args, self.state, self.control) - timers("interval-time").stop() torch.distributed.barrier() + self.control = self.callback_handler.on_save_begin(self.args, self.state, self.control) + self.train_engine.save_checkpoint( Path(self.args.output_dir), best_model_checkpoint=self.state.best_model_checkpoint, ) - timers("interval-time", log_level=0).start(barrier=True) - self.control = self.callback_handler.on_save(self.args, self.state, self.control) + if self.distributed_type == DistributedType.MEGATRON: + timers("interval-time", log_level=0).start(barrier=True) + def log(self, logs: Dict[str, float]) -> None: """ Log `logs` on the various objects watching training. @@ -828,6 +720,9 @@ def _load_rng_state(self, checkpoint): "\nThis won't yield the same results as if the training had not been interrupted." ) + def evaluate_dpo(self, step): + pass + def evaluate( self, eval_dataset: Optional[Dataset] = None, @@ -866,8 +761,13 @@ def evaluate( eval_iters = getattr(megatron_args, f"{eval_or_test}_iters", 0) if eval_iters > 0: max_eval_steps = eval_iters - else: + elif has_length(eval_dataloader): max_eval_steps = len(eval_dataloader) + else: + return None + + timers = get_timers() + timers(eval_or_test, log_level=0).start(barrier=True) else: # max_eval_steps = self.num_examples(eval_dataloader) raise ValueError(f"Evaluation on {self.distributed_type} not implement.") @@ -929,6 +829,11 @@ def evaluate( self.log(eval_log) + self.train_engine.train() + # TODO:(L1) place the following code to on_evaluate() + if self.distributed_type == DistributedType.MEGATRON: + timers(eval_or_test).stop() + timers.log([eval_or_test]) self.control = self.callback_handler.on_evaluate(self.args, self.state, self.control, eval_log) return eval_log @@ -947,6 +852,75 @@ def num_examples(self, dataloader: DataLoader) -> int: except (NameError, AttributeError, TypeError): # no dataset or length, estimate by length of dataloader return len(dataloader) * self.args.per_device_train_batch_size + def _get_train_num(self, args, train_dataloader, total_train_batch_size): + len_dataloader = None + if self.distributed_type == DistributedType.MEGATRON: + megatron_args = get_args() + if self.args.finetune_type is not None: + len_dataloader = len(train_dataloader) + num_update_steps_per_epoch = max(len_dataloader, 1) + num_examples = train_dataloader.num_examples + + # Internal checking + # args.max_steps has been calculated in MegatronEngine + assert args.max_steps > 0, "`max_steps` should be greater than 0" + assert ( + args.max_steps == megatron_args.train_iters + ), "Please ensure trainer args.max_steps is equal to megatron_args.train_iters." + max_steps = args.max_steps + num_train_epochs = args.max_steps // num_update_steps_per_epoch + int( + args.max_steps % num_update_steps_per_epoch > 0 + ) + # May be slightly incorrect if the last batch in the training dataloader has a smaller size but it's + # the best we can do. + num_train_samples = max_steps * total_train_batch_size + else: + if ( + args.max_steps > 0 + and megatron_args.train_iters is not None + and args.max_steps != megatron_args.train_iters + ): + logger.warning( + "args.max_steps will be overwritten by megatron_args.train_iters under MEGATRON training mode." + ) + args.max_steps = megatron_args.train_iters + max_steps = args.max_steps + num_train_epochs = sys.maxsize + num_update_steps_per_epoch = max_steps + num_examples = total_train_batch_size * max_steps + num_train_samples = max_steps * total_train_batch_size + elif has_length(train_dataloader): + len_dataloader = len(train_dataloader) + num_update_steps_per_epoch = len_dataloader // args.gradient_accumulation_steps + num_update_steps_per_epoch = max(num_update_steps_per_epoch, 1) + num_examples = self.num_examples(train_dataloader) + if args.max_steps > 0: + max_steps = args.max_steps + num_train_epochs = args.max_steps // num_update_steps_per_epoch + int( + args.max_steps % num_update_steps_per_epoch > 0 + ) + # May be slightly incorrect if the last batch in the training dataloader has a smaller size but it's + # the best we can do. + num_train_samples = args.max_steps * total_train_batch_size + else: + max_steps = math.ceil(args.num_train_epochs * num_update_steps_per_epoch) + num_train_epochs = math.ceil(args.num_train_epochs) + num_train_samples = self.num_examples(train_dataloader) * args.num_train_epochs # type: ignore[assignment] # noqa: E501 + elif args.max_steps > 0: # Rely on max_steps when dataloader does not have a working size + max_steps = args.max_steps + # Setting a very large number of epochs so we go as many times as necessary over the iterator. + num_train_epochs = sys.maxsize + num_update_steps_per_epoch = max_steps + num_examples = total_train_batch_size * args.max_steps + num_train_samples = args.max_steps * total_train_batch_size + else: + raise ValueError( + "args.max_steps must be set to a positive value if dataloader does not have a length, was" + f" {args.max_steps}" + ) + + return len_dataloader, max_steps, num_train_epochs, num_update_steps_per_epoch, num_examples, num_train_samples + def _nested_gather(self, tensors, name=None): """ Gather value of `tensors` (tensor or list/tuple of nested tensors) and convert them to numpy before diff --git a/atorch/trainer/debug_utils/debug_module.py b/atorch/trainer/debug_utils/debug_module.py new file mode 100644 index 0000000..6724184 --- /dev/null +++ b/atorch/trainer/debug_utils/debug_module.py @@ -0,0 +1,132 @@ +import inspect +import os + +import torch + +from atorch.common.log_utils import default_logger as logger +from atorch.trainer.args import AtorchTrainingArgs +from atorch.trainer.trainer_callback import AtorchTrainerCallback, AtorchTrainerControl, AtorchTrainerState +from atorch.utils.import_util import is_megatron_lm_available + +if is_megatron_lm_available(): + from megatron.core.package_info import __version__ as megatron_version + from megatron.training import get_args + + from atorch.trainer.megatron import AtorchMegatronEngine + + +class DebugCallback(AtorchTrainerCallback): + def __init__(self, train_engine): + """ + Should be called after initializing Megatron. + """ + super().__init__() + self.train_engine: AtorchMegatronEngine = train_engine + self.module_list = self.train_engine.module + self.modules = self.train_engine.module[0].modules() + self.args = get_args() + + self.activation_output_dir = os.path.join(self.args.save, f"activation_{megatron_version}") + if self.args.rank == 0: + os.makedirs(self.activation_output_dir, exist_ok=True) + + self.module_to_name = {} + if self.args.rank == 0: + for name, sub_module in self.train_engine.module[0].named_modules(): + m = inspect.getmodule(sub_module.__class__) + if m is not None: + module_name = f"{m.__name__}.{sub_module._get_name()}" + else: + module_name = f"{sub_module._get_name()}" + logger.info(f"[Rank {self.args.rank}] module: {name} module_name: {module_name}") + self.module_to_name[sub_module] = name + + def register_forward_pre_hook(self): + def _hook(module: torch.nn.Module, args, kwargs): + module_name = module._get_name() + if self.args.rank == 0 and self.args.global_step == 0: + logger.info( + f"[Rank {self.args.rank}] fwd pre hook: {module_name} len(args) {len(args)} kwargs.keys() {kwargs.keys()}" # noqa: E501 + ) + if module_name == "BailingMoeModel": + info = "" + for i, t in enumerate(args): + info += f" arg{i}: {t.shape if isinstance(t, torch.Tensor) else type(t)}" + logger.info(f"[Rank {self.args.rank}] fwd pre hook: {module_name} {info}") # noqa: E501 + + return _hook + + def register_forward_post_hook(self): + def _save(obj, name): + path = os.path.join( + self.activation_output_dir, + f"step{self.args.global_step}_micro{self.args.micro_step}_rank{self.args.rank}_{name}.pth", + ) + logger.info(f"[Rank {self.args.rank}] Saving activation to {path}") + torch.save(obj, path) + + def _hook(module: torch.nn.Module, args, kwargs, result): + module_name = module._get_name() + if self.args.rank == 0 and self.args.global_step == 0: + result_to_print = ( + result.keys() + if isinstance(result, dict) + else result.shape + if isinstance(result, torch.Tensor) + else type(result) + ) + logger.info( + f"[Rank {self.args.rank}] fwd post hook: {module_name} len(args) {len(args)} kwargs.keys() {kwargs.keys()} result: {result_to_print}" # noqa: E501 + ) + if module_name == "BailingMoeModel": + info = "" + for i, t in enumerate(args): + info += f" arg{i}: {t.shape if isinstance(t, torch.Tensor) else type(t)}" + logger.info(f"[Rank {self.args.rank}] fwd post hook: {module_name} {info}") # noqa: E501 + if self.args.rank == 0 and self.args.global_step in [2, 3, 4]: + name: str = self.module_to_name[module] + if "layers" in name: + splits = name.split(".") + layer_id = int(splits[4]) if len(splits) >= 5 else -1 # noqa: F841 + if ( + module_name + in [ + "SelfAttention", + "MoELayer", + "TopKRouter", + "TEGroupedMLP", + "SharedExpertMLP", + "TransformerLayer", + ] + or "mlp" in name + or ( + "self_attention" in name + and module_name not in ["UnfusedDotProductAttention", "FusedScaleMaskSoftmax", "Dropout"] + ) + ): + _save(result, name) + elif ( + module_name not in ["DistributedDataParallel", "Float16Module"] + and "output_layer" not in module_name + ): + _save(result, name) + if "final_layernorm" in name: + _save(args, name + "_input_args") + _save(kwargs, name + "_input_kwargs") + + return _hook + + def on_train_begin( + self, + atorch_training_args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + """ + Event called at the beginning of training, after Megatron initialization and building dataloader. + """ + for m in self.modules: + if isinstance(m, torch.nn.Module): + # m.register_forward_pre_hook(self.register_forward_pre_hook(), with_kwargs=True) + m.register_forward_hook(self.register_forward_post_hook(), with_kwargs=True) diff --git a/atorch/trainer/megatron/__init__.py b/atorch/trainer/megatron/__init__.py index f310c39..e855e08 100644 --- a/atorch/trainer/megatron/__init__.py +++ b/atorch/trainer/megatron/__init__.py @@ -1,3 +1,3 @@ -from .megatron_dataloader import AtorchMegatronDataloader +from .megatron_dataloader import AtorchMegatronDataloader, skip_first_batches_for_megatron_dataloader from .megatron_train_step import BertTrainStep, GPTTrainStep, MegatronTrainStep, T5TrainStep -from .megatron_wrapper import AtorchMegatronEngine +from .megatron_wrapper import AtorchMegatronEngine, MegatronCallback diff --git a/atorch/trainer/megatron/megatron_async_save.py b/atorch/trainer/megatron/megatron_async_save.py index fbcc9bf..ad75e4a 100644 --- a/atorch/trainer/megatron/megatron_async_save.py +++ b/atorch/trainer/megatron/megatron_async_save.py @@ -89,6 +89,7 @@ def save( # type: ignore[override] optimizer=None, scheduler=None, num_floating_point_operations_so_far=None, + **kwargs, ): megatron_args = get_args() async_timeout = train_args.flash_checkpoint_timeout diff --git a/atorch/trainer/megatron/megatron_ckpt_loader.py b/atorch/trainer/megatron/megatron_ckpt_loader.py index 41b87d2..d1b9d83 100644 --- a/atorch/trainer/megatron/megatron_ckpt_loader.py +++ b/atorch/trainer/megatron/megatron_ckpt_loader.py @@ -18,6 +18,7 @@ def load( # type: ignore[override] optimizer=None, scheduler=None, train_args: AtorchTrainingArgs = None, + **kwargs, ) -> Tuple[int, int]: pass @@ -30,6 +31,7 @@ def load( # type: ignore[override] optimizer=None, scheduler=None, train_args: AtorchTrainingArgs = None, + **kwargs, ): assert model is not None, "Megatron load model should not be None" assert optimizer is not None, "Megatron load optimizer should not be None" @@ -46,7 +48,9 @@ def load( # type: ignore[override] from megatron.training.checkpointing import load_checkpoint - iteration, num_floating_point_operations_so_far = load_checkpoint(model, optimizer, scheduler) + iteration, num_floating_point_operations_so_far = load_checkpoint( + model, optimizer, scheduler, strict=train_args.resume_strict + ) # pragma: no cover torch.distributed.barrier() diff --git a/atorch/trainer/megatron/megatron_ckpt_saver.py b/atorch/trainer/megatron/megatron_ckpt_saver.py index 3b8daa0..0921f36 100644 --- a/atorch/trainer/megatron/megatron_ckpt_saver.py +++ b/atorch/trainer/megatron/megatron_ckpt_saver.py @@ -2,6 +2,7 @@ import shutil from abc import ABC, abstractmethod from pathlib import Path +from typing import Any, Callable, Dict, Optional import torch from megatron.training import get_args @@ -58,11 +59,15 @@ def save( # type: ignore[override] optimizer=None, scheduler=None, num_floating_point_operations_so_far=None, + **kwargs, ): pass class MegatronOriginSaver(MegatronCkptSaver): + + _checkpointing_context: Optional[Dict] = None + def save( # type: ignore[override] self, iteration: int, @@ -73,6 +78,7 @@ def save( # type: ignore[override] optimizer=None, scheduler=None, num_floating_point_operations_so_far=None, + **kwargs, ): megatron_args = get_args() @@ -80,6 +86,9 @@ def save( # type: ignore[override] if output_dir is not None: megatron_args.save = output_dir + if megatron_args.ckpt_assume_constant_structure and MegatronOriginSaver._checkpointing_context is None: + MegatronOriginSaver._checkpointing_context = {} + if is_main_process(): checkpoint_dir_path = self.get_interation_path(output_dir, iteration, return_base_dir=True) os.makedirs(checkpoint_dir_path, exist_ok=True) @@ -91,12 +100,22 @@ def save( # type: ignore[override] from megatron.training.checkpointing import save_checkpoint + extra_args: Dict[str, Any] = {} + + custom_async_finalize_fn: Callable = ( + kwargs["custom_async_finalize_fn"] if "custom_async_finalize_fn" in kwargs else None + ) + if custom_async_finalize_fn is not None: + extra_args.update(custom_async_finalize_fn=custom_async_finalize_fn) + save_checkpoint( iteration, module, optimizer, scheduler, num_floating_point_operations_so_far=num_floating_point_operations_so_far, + checkpointing_context=MegatronOriginSaver._checkpointing_context, + **extra_args, ) torch.distributed.barrier() diff --git a/atorch/trainer/megatron/megatron_dataloader.py b/atorch/trainer/megatron/megatron_dataloader.py index 326458c..23db4f0 100644 --- a/atorch/trainer/megatron/megatron_dataloader.py +++ b/atorch/trainer/megatron/megatron_dataloader.py @@ -2,7 +2,7 @@ from typing import Callable, Iterator, List, Union import torch -from torch.utils.data import DataLoader +from torch.utils.data import BatchSampler, DataLoader, IterableDataset from atorch.common.log_utils import default_logger as logger from atorch.trainer.base.dataloader import AtorchDataloader @@ -29,6 +29,7 @@ "generator": None, "prefetch_factor": 2, "persistent_workers": False, + "pin_memory_device": "", } @@ -349,172 +350,193 @@ def __len__(self): return data_iterator -def wrap_megatron_dataloader( - data_iterator: Union[Iterator, List[Iterator], DataLoader, List[DataLoader]], - dataset_type: str, # ["train", "eval", "test"] - is_post_training: bool, -): +class MegatronIteratorWrapper: """ - Args: - data_iterator: - Returns: + MegatronIteratorWrapper is used in pretrain, just one epoch. + """ + + def __init__( + self, + data_iterator: Union[Iterator, List[Iterator]], + ) -> None: + self.data_iterator = data_iterator + + def __iter__(self): + return self + + def __next__(self): + return self.data_iterator + +class MegatronDataloaderWrapper: + """ + MegatronDataloaderWrapper is used in post-training scene, supporting multi-epochs. """ - args = get_args() + def __init__( + self, + dataloader: Union[DataLoader, List[DataLoader]], + is_post_training, + ) -> None: + self.dataloader = dataloader + + # The follow attributes are None if self.dataloader is None. + self.dataset = None + self.sampler = None + self.batch_sampler = None + + dataloader_length = 0 + dataset_length = 0 + batch_size = 0 + drop_last = 1 + + def _extract_info_from_dataloader(d): + assert hasattr(d, "dataset"), "dataloader must has 'dataset' attribute." + assert hasattr(d, "sampler"), "dataloader must has 'sampler' attribute." + assert hasattr(d, "batch_sampler"), "dataloader must has 'batch_sampler' attribute." + self.dataset = d.dataset + self.sampler = d.sampler + self.batch_sampler = d.batch_sampler + assert hasattr(self.batch_sampler, "drop_last") + assert hasattr(self.batch_sampler, "batch_size") or hasattr(self.batch_sampler, "micro_batch_size") + _drop_last = int(self.batch_sampler.drop_last) + _batch_size = ( + getattr(self.batch_sampler, "batch_size", None) + or getattr(self.batch_sampler, "micro_batch_size", None) + or 0 + ) + _dataloader_length = len(d) + _dataset_length = len(self.dataset) - class MegatronIteratorWrapper: - """ - MegatronIteratorWrapper is used in pretrain, just one epoch. - """ + return _dataloader_length, _dataset_length, _batch_size, _drop_last - def __init__( - self, - data_iterator: Union[Iterator, List[Iterator]], - ) -> None: - self.data_iterator = data_iterator + args = get_args() - def __iter__(self): - return self + if args.virtual_pipeline_model_parallel_size is not None: + assert isinstance(self.dataloader, list), "VPP requires dataloader to be a list." + for d in self.dataloader: + if d is not None: + dataloader_length, dataset_length, batch_size, drop_last = _extract_info_from_dataloader(d) + elif self.dataloader is not None: + dataloader_length, dataset_length, batch_size, drop_last = _extract_info_from_dataloader(self.dataloader) - def __next__(self): - return self.data_iterator + size_info = torch.tensor( + [dataloader_length, dataset_length, batch_size, drop_last], + dtype=torch.int64, + device=torch.cuda.current_device(), + ) + # assert: dataloader on model_parallel_rank 0 must be a real dataloader instance, not be None. + model_parallel_src_rank = torch.distributed.get_process_group_ranks(group=mpu.get_model_parallel_group())[0] + torch.distributed.broadcast(size_info, model_parallel_src_rank, group=mpu.get_model_parallel_group()) + + # The follow three size info will be set to the value on model_parallel rank 0 + # via broadcast in model_parallel group. + self.dataloader_length = size_info[0].item() + self.dataset_length = size_info[1].item() + self.dataloader_batch_size = size_info[2].item() + self.drop_last = bool(size_info[3].item()) + + # Check batch_size in dataloader + if not is_post_training: + assert ( + self.dataloader_batch_size == args.micro_batch_size + ), "dataloader's batch_size should be micro_batch_size when using Megatron." + num_microbatches = args.global_batch_size // (args.micro_batch_size * args.data_parallel_size) + + if self.drop_last: + self.num_update_steps_per_epoch = self.dataloader_length // num_microbatches + else: + self.num_update_steps_per_epoch = math.ceil(self.dataloader_length / num_microbatches) + + if self.dataloader is not None: + rank = torch.distributed.get_rank() + if self.dataloader_length == self.dataset_length: + logger.warning( + f"WARNING [Rank {rank}] dataloader and dataset has the same length {self.dataloader_length}, " + "unexpected!" + ) + logger.info( + f"[Rank {rank}] dataset_length {self.dataset_length} dataloader_length {self.dataloader_length}" + f" batch_sampler_length {len(self.batch_sampler) if self.batch_sampler is not None else 0}" + f" batch_size {self.dataloader_batch_size} num_microbatches {num_microbatches}" + f" num_update_steps_per_epoch {self.num_update_steps_per_epoch} drop_last {self.drop_last}", + ) + + self.skip_batches = 0 + + self._reset() + + def _reset(self): + self._num_yielded = self.skip_batches + if isinstance(self.dataloader, list): + self._dataloader_iter = [iter(d) if d is not None else None for d in self.dataloader] + else: + self._dataloader_iter = iter(self.dataloader) if self.dataloader is not None else None + + def __iter__(self): + self._reset() + + logger.info(f" [Rank {torch.distributed.get_rank()}] Reset dataloader, skip first {self._num_yielded} batches.") - class MegatronDataloaderWrapper: + # Skipping first batches only occurs in the first training epoch. + if self.skip_batches > 0: + self.skip_batches = 0 + return self + + def __next__(self): + if self._num_yielded < self.num_update_steps_per_epoch: + self._num_yielded += 1 + return self._dataloader_iter + else: + raise StopIteration + + def __len__(self): + return self.num_update_steps_per_epoch + + def set_epoch(self, epoch): """ - MegatronDataloaderWrapper is used in post-training scene, supporting multi-epochs. + Call self.sampler.set_epoch() if it is not None, only used in finetune scene. """ - def __init__( - self, - dataloader: Union[DataLoader, List[DataLoader]], - ) -> None: - self.dataloader = dataloader - - # The follow attributes are None if self.dataloader is None. - self.dataset = None - self.sampler = None - self.batch_sampler = None - - dataloader_length = 0 - dataset_length = 0 - batch_size = 0 - drop_last = 1 - - def _extract_info_from_dataloader(d): - assert hasattr(d, "dataset"), "dataloader must has 'dataset' attribute." - assert hasattr(d, "sampler"), "dataloader must has 'sampler' attribute." - assert hasattr(d, "batch_sampler"), "dataloader must has 'batch_sampler' attribute." - self.dataset = d.dataset - self.sampler = d.sampler - self.batch_sampler = d.batch_sampler - assert hasattr(self.batch_sampler, "drop_last") - assert hasattr(self.batch_sampler, "batch_size") or hasattr(self.batch_sampler, "micro_batch_size") - _drop_last = int(self.batch_sampler.drop_last) - _batch_size = ( - getattr(self.batch_sampler, "batch_size", None) - or getattr(self.batch_sampler, "micro_batch_size", None) - or 0 - ) - _dataloader_length = len(d) - _dataset_length = len(self.dataset) - - return _dataloader_length, _dataset_length, _batch_size, _drop_last - - if args.virtual_pipeline_model_parallel_size is not None: - assert isinstance(self.dataloader, list), "VPP requires dataloader to be a list." - for d in self.dataloader: - if d is not None: - dataloader_length, dataset_length, batch_size, drop_last = _extract_info_from_dataloader(d) - elif self.dataloader is not None: - dataloader_length, dataset_length, batch_size, drop_last = _extract_info_from_dataloader( - self.dataloader - ) - - size_info = torch.tensor( - [dataloader_length, dataset_length, batch_size, drop_last], - dtype=torch.int64, - device=torch.cuda.current_device(), - ) - # assert: dataloader on model_parallel_rank 0 must be a real dataloader instance, not be None. - model_parallel_src_rank = torch.distributed.get_process_group_ranks(group=mpu.get_model_parallel_group())[0] - torch.distributed.broadcast(size_info, model_parallel_src_rank, group=mpu.get_model_parallel_group()) - - # The follow three size info will be set to the value on model_parallel rank 0 - # via broadcast in model_parallel group. - self.dataloader_length = size_info[0].item() - self.dataset_length = size_info[1].item() - self.dataloader_batch_size = size_info[2].item() - self.drop_last = bool(size_info[3].item()) - - # Check batch_size in dataloader - if not is_post_training: - assert ( - self.dataloader_batch_size == args.micro_batch_size - ), "dataloader's batch_size should be micro_batch_size when using Megatron." - num_microbatches = args.global_batch_size // (args.micro_batch_size * args.data_parallel_size) - - if self.drop_last: - self.num_update_steps_per_epoch = self.dataloader_length // num_microbatches - else: - self.num_update_steps_per_epoch = math.ceil(self.dataloader_length / num_microbatches) - - if self.dataloader is not None: - rank = torch.distributed.get_rank() - if self.dataloader_length == self.dataset_length: - logger.warning( - f"WARNING [Rank {rank}] dataloader and dataset has the same length {self.dataloader_length}, " - "unexpected!" - ) - logger.info( - f"[Rank {rank}] dataset_length {self.dataset_length} dataloader_length {self.dataloader_length}" - f" batch_sampler_length {len(self.batch_sampler) if self.batch_sampler is not None else 0}" - f" batch_size {self.dataloader_batch_size} num_microbatches {num_microbatches}" - f" num_update_steps_per_epoch {self.num_update_steps_per_epoch} drop_last {self.drop_last}", - ) + if self.sampler is not None or self.batch_sampler is not None: + set_epoch = getattr(self.sampler, "set_epoch", None) or getattr(self.batch_sampler, "set_epoch", None) + assert set_epoch is not None and isinstance( + set_epoch, Callable + ), "dataloader's sampler or batch_sampler should have set_epoch() method when do post-training." + logger.info(f"Calling set_epoch({epoch})") + set_epoch(epoch) - self._reset() + @property + def num_examples(self): + """ + Return the num of examples in dataset. self.dataset_length is valid on each rank, + because it received correct value from model parallel rank 0 via broadcasting. + """ + return self.dataset_length - def _reset(self): - self._num_yielded = 0 - if isinstance(self.dataloader, list): - self._dataloader_iter = [iter(d) if d is not None else None for d in self.dataloader] - else: - self._dataloader_iter = iter(self.dataloader) if self.dataloader is not None else None + def set_skip_batches(self, skip_batches): + """ + Set how many batches should be skipped to support resuming training. + """ + assert skip_batches >= 0, f"skip_batches should be >= 0, but got {skip_batches}." + if skip_batches >= self.num_update_steps_per_epoch: + raise ValueError( + f"skip_batches={skip_batches} >= num_update_steps_per_epoch={self.num_update_steps_per_epoch}, not reasonable!" # noqa: E501 + ) + self.skip_batches = skip_batches - def __iter__(self): - self._reset() - return self - def __next__(self): - if self._num_yielded < self.num_update_steps_per_epoch: - self._num_yielded += 1 - return self._dataloader_iter - else: - raise StopIteration +def wrap_megatron_dataloader( + data_iterator: Union[Iterator, List[Iterator], DataLoader, List[DataLoader]], + dataset_type: str, # ["train", "eval", "test"] + is_post_training: bool, +): + """ + Args: + data_iterator: + Returns: - def __len__(self): - return self.num_update_steps_per_epoch - - def set_epoch(self, epoch): - """ - Call self.sampler.set_epoch() if it is not None, only used in finetune scene. - """ - - if self.sampler is not None: - assert hasattr(self.sampler, "set_epoch") and isinstance( - self.sampler.set_epoch, Callable - ), "dataloader's sampler should have set_epoch() method when finetune." - logger.info(f"Calling set_epoch({epoch})") - self.sampler.set_epoch(epoch) - - @property - def num_examples(self): - """ - Return the num of examples in dataset. self.dataset_length is valid on each rank, - because it received correct value from model parallel rank 0 via broadcasting. - """ - return self.dataset_length + """ def _check_dataloader_type(dataloader, dataloader_type, pre_info): if isinstance(dataloader, list): @@ -534,7 +556,7 @@ def _check_dataloader_type(dataloader, dataloader_type, pre_info): # torch.util.data.Dataloader type is required for dataloader object in finetune scene. _check_dataloader_type(data_iterator, DataLoader, "post-train") - return MegatronDataloaderWrapper(data_iterator) + return MegatronDataloaderWrapper(data_iterator, is_post_training) else: ####### Compat old code in antllm. To be removed rank = torch.distributed.get_rank() @@ -570,7 +592,7 @@ def _check_dataloader_type(dataloader, dataloader_type, pre_info): _check_dataloader_type(data_iterator, DataLoader, "pretrain") if dataset_type == "test": - return MegatronDataloaderWrapper(data_iterator) + return MegatronDataloaderWrapper(data_iterator, is_post_training) else: if data_iterator is not None: if isinstance(data_iterator, list): # for VPP @@ -582,3 +604,109 @@ def _check_dataloader_type(dataloader, dataloader_type, pre_info): raise ValueError(f"Unexpected iterator_type {iterator_type}, please check the broadcast of iterator_type.") return MegatronIteratorWrapper(data_iterator) + + +class SkipBatchSampler(BatchSampler): + """ + A `torch.utils.data.BatchSampler` that skips the first `n` batches of another `torch.utils.data.BatchSampler`. + Should not be used if the original dataloader is a `StatefulDataLoader`. + """ + + def __init__(self, batch_sampler, skip_batches=0): + self.batch_sampler = batch_sampler + self.batch_size = batch_sampler.batch_size + self.drop_last = batch_sampler.drop_last + self.skip_batches = skip_batches + + def __iter__(self): + for index, samples in enumerate(self.batch_sampler): + if index >= self.skip_batches: + yield samples + + def __len__(self) -> int: + return len(self.batch_sampler) + + +def skip_first_batches_for_megatron_dataloader(megatron_dataloader, num_batches=0): + """ + Creates a `torch.utils.data.DataLoader` that will efficiently skip the first `num_batches`. Should not be used if + the original dataloader is a `StatefulDataLoader`. + """ + args = get_args() + + assert isinstance( + megatron_dataloader, MegatronDataloaderWrapper + ), f"Only support dataloader with MegatronDataloaderWrapper type, but got {type(megatron_dataloader)}." + + dataloader = megatron_dataloader.dataloader + + if args.virtual_pipeline_model_parallel_size: + """ + assume pp size is 4, vpp size is 3, dataloader on pp stages will be: + + pp_rank + 0: [DadaLoader(), None, None] + 1: [None, None, None] + 2: [None, None, None] + 3: [None, None, DadaLoader()] + """ + assert isinstance(dataloader, List), "VPP requires dataloader to be a list." + + dataloader_position = -1 + for i, d in enumerate(dataloader): + if d is not None: + dataloader_position = i + dataloader = d + break + if dataloader_position == -1: + dataloader = None + + if dataloader is None: + if args.virtual_pipeline_model_parallel_size: + new_dataloader = [None for _ in range(args.virtual_pipeline_model_parallel_size)] + else: + new_dataloader = None + else: + dataset = dataloader.dataset + sampler_is_batch_sampler = False + assert not isinstance( + dataset, IterableDataset + ), f"dataset {type(dataset)} is IterableDataset class, not supported resuming training in Megatron. Please contact ATorch developer." # noqa: E501 + + sampler_is_batch_sampler = isinstance(dataloader.sampler, BatchSampler) + batch_sampler = dataloader.sampler if sampler_is_batch_sampler else dataloader.batch_sampler + + assert args.rampup_batch_size is None, "rampup_batch_size is not supported in post-training." + num_microbatches = args.global_batch_size // (args.micro_batch_size * args.data_parallel_size) + skip_microbatches = num_batches * num_microbatches + + new_batch_sampler = SkipBatchSampler(batch_sampler, skip_batches=skip_microbatches) + + # We ignore all of those since they are all dealt with by our new_batch_sampler + ignore_kwargs = [ + "batch_size", + "shuffle", + "sampler", + "batch_sampler", + "drop_last", + ] + + kwargs = { + k: getattr(dataloader, k, _PYTORCH_DATALOADER_KWARGS[k]) + for k in _PYTORCH_DATALOADER_KWARGS + if k not in ignore_kwargs + } + + skip_batches_dataloader = DataLoader(dataset, batch_sampler=new_batch_sampler, **kwargs) + + if args.virtual_pipeline_model_parallel_size: + new_dataloader = [None for _ in range(args.virtual_pipeline_model_parallel_size)] + new_dataloader[dataloader_position] = skip_batches_dataloader + else: + new_dataloader = skip_batches_dataloader + + megatron_dataloader = MegatronDataloaderWrapper(new_dataloader, is_post_training=True) + + megatron_dataloader.set_skip_batches(num_batches) + + return megatron_dataloader diff --git a/atorch/trainer/megatron/megatron_wrapper.py b/atorch/trainer/megatron/megatron_wrapper.py index 23758e4..f5af880 100644 --- a/atorch/trainer/megatron/megatron_wrapper.py +++ b/atorch/trainer/megatron/megatron_wrapper.py @@ -2,9 +2,12 @@ Megatron wrapper. """ import dataclasses +import gc +import inspect import math import os import random +import re import sys from pathlib import Path from typing import Dict, List, Optional, Union @@ -13,11 +16,13 @@ import torch import torch.distributed from dependency_injector.wiring import Provide, inject +from packaging.version import Version from torch.utils.data import DataLoader from transformers.trainer import TRAINER_STATE_NAME from atorch.common.log_utils import default_logger as logger -from atorch.trainer.args import AtorchTrainingArgs +from atorch.common.log_utils import log_rank_0 +from atorch.trainer.args import AtorchTrainingArgs, MegatronArgs from atorch.trainer.base.atorch_container import AtorchTrainerContainer from atorch.trainer.base.atorch_train_engine import AtorchTrainEngine from atorch.trainer.megatron.megatron_ckpt_loader import MegatronCkptLoader @@ -28,7 +33,7 @@ wrap_megatron_dataloader, ) from atorch.trainer.megatron.megatron_train_step import BertTrainStep, GPTTrainStep, MegatronTrainStep, T5TrainStep -from atorch.trainer.trainer_callback import AtorchTrainerState +from atorch.trainer.trainer_callback import AtorchTrainerCallback, AtorchTrainerControl, AtorchTrainerState from atorch.trainer.utils import ( broadcast_spike_loss_ratio_in_pp_group, calc_params_std, @@ -38,7 +43,7 @@ training_log, ) from atorch.utils.import_util import is_megatron_lm_available, is_torch_npu_available -from atorch.utils.version import is_megatron_version_bigger_than +from atorch.utils.version import get_megatron_version, is_megatron_version_bigger_than from atorch.utils.virtual_optimizer.megatron_virtual_optimizer import get_megatron_virtual_optimizer if is_megatron_lm_available(): @@ -46,20 +51,26 @@ from megatron.core.distributed import DistributedDataParallel as MegatronDDP from megatron.core.distributed import finalize_model_grads from megatron.core.enums import ModelType - from megatron.core.optimizer import MegatronOptimizer, OptimizerConfig + from megatron.core.optimizer import MegatronOptimizer, OptimizerConfig, get_megatron_optimizer from megatron.core.pipeline_parallel import get_forward_backward_func + from megatron.core.transformer.moe.moe_layer import MoELayer from megatron.core.utils import get_model_config from megatron.legacy.model import BertModel, GPTModel, T5Model from megatron.legacy.model.classification import Classification from megatron.legacy.model.module import MegatronModule - from megatron.training import get_args, get_tensorboard_writer, print_rank_last + from megatron.training import get_args, get_tensorboard_writer, initialize_megatron, print_rank_last try: from megatron.training import get_num_microbatches, update_num_microbatches except ImportError: from megatron.core.num_microbatches_calculator import get_num_microbatches, update_num_microbatches from megatron.training.arguments import core_transformer_config_from_args, parse_args, validate_args - from megatron.training.checkpointing import load_args_from_checkpoint + from megatron.training.checkpointing import ( + checkpoint_exists, + load_args_from_checkpoint, + load_checkpoint, + save_checkpoint, + ) from megatron.training.global_vars import get_timers, set_global_variables from megatron.training.initialize import ( _compile_dependencies, @@ -80,22 +91,75 @@ get_optimizer_param_scheduler, num_floating_point_operations, ) - from megatron.training.utils import calc_params_l2_norm, print_rank_0, unwrap_model + from megatron.training.utils import calc_params_l2_norm, unwrap_model from megatron.training.yaml_arguments import validate_yaml DATALOADER_INDEX_MAPPER = dict(train=0, eval=1, test=2) -def initialize_megatron( +def setup_args(extra_args_provider=None, args_defaults={}, ignore_unknown_args=False): + # Parse arguments + args = parse_args(extra_args_provider, ignore_unknown_args) + + # Set defaults + for key, value in args_defaults.items(): + if getattr(args, key, None) is not None: + if args.rank == 0: + print( + f"WARNING: overriding default arguments for " f"{key}:{getattr(args, key)} with {key}:{value}", + flush=True, + ) + # set extra_configs. + setattr(args, key, value) + + return args + + +def set_deterministic_algorithms(args): + """ + args: Megatron args, acquired by get_args() + """ + # Megatron's `deterministic_mode` arg is introduced after core_0.8.0 version. + if is_megatron_version_bigger_than("0.8.0"): + args.deterministic_mode = True + args.use_flash_attn = False + args.cross_entropy_loss_fusion = False + + if not is_torch_npu_available(): + # On GPU env, bias_dropout_fusion will effect the accuracy of loss and grad when resuming + # training from a checkpoint. So if you want to use deterministic algorithms, set it to False. + args.bias_dropout_fusion = False + + # Set env variables about deterministic mode + if os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1") != "0": + logger.info("For deterministic algo, env [NVTE_ALLOW_NONDETERMINISTIC_ALGO] will be set to '0'.") + os.environ["NVTE_ALLOW_NONDETERMINISTIC_ALGO"] = "0" + + all_reduce_choices = ["Tree", "Ring", "CollnetDirect", "CollnetChain", "^NVLS"] + if os.getenv("NCCL_ALGO") not in all_reduce_choices: + logger.info("For deterministic algo, env [NCCL_ALGO] will be set to 'Ring'.") + os.environ["NCCL_ALGO"] = "Ring" + + cublas_workspace_config_choices = [":4096:8", ":16:8"] + if os.getenv("CUBLAS_WORKSPACE_CONFIG") not in cublas_workspace_config_choices: + logger.info("For deterministic algo, env [CUBLAS_WORKSPACE_CONFIG] will be set to ':4096:8'.") + os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" + + torch.use_deterministic_algorithms(True, warn_only=True) + + +# Deprecated! initialize_megatron_legacy() is deprecated when megatron >= 0.12 +def initialize_megatron_legacy( extra_args_provider=None, args_defaults={}, ignore_unknown_args=False, allow_no_cuda=False, skip_mpu_initialization=False, - use_deterministic_algorithms=False, get_embedding_ranks=None, get_position_embedding_ranks=None, + get_output_layer_ranks=None, + use_deterministic_algorithms=False, ): """Set global variables, initialize distributed, and set autoresume and random seeds. @@ -105,55 +169,27 @@ def initialize_megatron( Returns a function to finalize distributed env initialization (optionally, only when args.lazy_mpu_init == True) """ + + assert not is_megatron_version_bigger_than( + "0.12.0" + ), "Please call original initialize_megatron() under Megatron 0.12!" + if not allow_no_cuda: # Make sure cuda is available. assert torch.cuda.is_available(), "Megatron requires CUDA." - # Parse arguments - args = parse_args(extra_args_provider, ignore_unknown_args) - # Set defaults - for key, value in args_defaults.items(): - if getattr(args, key, None) is not None: - if args.rank == 0: - print( - f"WARNING: overriding default arguments for " f"{key}:{getattr(args, key)} with {key}:{value}", - flush=True, - ) - # TODO: extra key. - setattr(args, key, value) + args = setup_args(extra_args_provider, args_defaults, ignore_unknown_args) if use_deterministic_algorithms: - # Megatron's `deterministic_mode` arg is introduced after core_0.8.0 version. - if is_megatron_version_bigger_than("0.8.0"): - args.deterministic_mode = True - args.use_flash_attn = False - args.cross_entropy_loss_fusion = False - - if not is_torch_npu_available(): - # On GPU env, bias_dropout_fusion will effect the accuracy of loss and grad when resuming - # training from a checkpoint. So if you want to use deterministic algorithms, set it to False. - args.bias_dropout_fusion = False - - # Set env variables about deterministic mode - if os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1") != "0": - logger.info("For deterministic algo, env [NVTE_ALLOW_NONDETERMINISTIC_ALGO] will be set to '0'.") - os.environ["NVTE_ALLOW_NONDETERMINISTIC_ALGO"] = "0" - - all_reduce_choices = ["Tree", "Ring", "CollnetDirect", "CollnetChain", "^NVLS"] - if os.getenv("NCCL_ALGO") not in all_reduce_choices: - logger.info("For deterministic algo, env [NCCL_ALGO] will be set to 'Ring'.") - os.environ["NCCL_ALGO"] = "Ring" - - cublas_workspace_config_choices = [":4096:8", ":16:8"] - if os.getenv("CUBLAS_WORKSPACE_CONFIG") not in cublas_workspace_config_choices: - logger.info("For deterministic algo, env [CUBLAS_WORKSPACE_CONFIG] will be set to ':4096:8'.") - os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" - - device_count = torch.cuda.device_count() - args.local_rank = torch.distributed.get_rank() % device_count + set_deterministic_algorithms(args) if args.use_checkpoint_args or args_defaults.get("use_checkpoint_args", False): - assert args.load is not None, "--use-checkpoints-args requires --load argument" + assert args.load is not None, "--use-checkpoint-args requires --load argument" + assert getattr(args, "non_persistent_ckpt_type", None) != "local", ( + "--use-checkpoint-args is not supported with --non_persistent_ckpt_type=local. " + "Two-stage checkpoint loading is not implemented, and all arguments must be defined " + "before initializing LocalCheckpointManager." + ) load_args_from_checkpoint(args) if args.yaml_cfg is not None: @@ -165,24 +201,32 @@ def initialize_megatron( # tensorboard-writer, and timers. set_global_variables(args) + # set logging level + if is_megatron_version_bigger_than("0.8.0"): + from megatron.training.initialize import setup_logging + + setup_logging() + # torch.distributed initialization def finish_mpu_init(): args = get_args() # Pytorch distributed. - import inspect - if "get_embedding_ranks" not in inspect.signature(_initialize_distributed).parameters: + if not is_megatron_version_bigger_than("0.9.0"): _initialize_distributed() + elif "get_output_layer_ranks" in inspect.signature(_initialize_distributed).parameters: + _initialize_distributed(get_embedding_ranks, get_position_embedding_ranks, get_output_layer_ranks) else: _initialize_distributed(get_embedding_ranks, get_position_embedding_ranks) # Random seeds for reproducibility. if args.rank == 0: print("> setting random seeds to {} ...".format(args.seed)) - _set_random_seed(args.seed, args.data_parallel_random_init) - if use_deterministic_algorithms: - torch.use_deterministic_algorithms(True) + if is_megatron_version_bigger_than("0.11.0"): + _set_random_seed(args.seed, args.data_parallel_random_init, args.te_rng_tracker, args.inference_rng_tracker) + else: + _set_random_seed(args.seed, args.data_parallel_random_init) if skip_mpu_initialization: return None @@ -215,8 +259,9 @@ def finish_mpu_init(): return None -def prepare_optimizer(train_args, args, megatron_args, model): +def prepare_optimizer(train_args: AtorchTrainingArgs, megatron_args: MegatronArgs, model): logger.info("Preparing optimizer") + args = get_args() kwargs = {} for f in dataclasses.fields(OptimizerConfig): if hasattr(args, f.name): @@ -225,8 +270,11 @@ def prepare_optimizer(train_args, args, megatron_args, model): config.timers = get_timers() # NOTE use_local_sgd should be injected with local_sgd_arg_provider if hasattr(args, "use_local_sgd") and args.use_local_sgd: + if get_megatron_version() != Version("0.9.0"): + logger.warning("[WARNING!!!!!!] local sgd is only tested under Megatron 0.9.0 !!!") + from atorch.local_sgd.configs import GTAConfig, LocalSGDConfig, OuterOptimizerConfig - from atorch.local_sgd.megatron import get_megatron_optimizer + from atorch.local_sgd.megatron import get_megatron_optimizer as get_megatron_optimizer_local_sgd local_sgd_config = LocalSGDConfig( local_sgd_sync_interval=args.local_sgd_sync_interval, @@ -256,7 +304,7 @@ def prepare_optimizer(train_args, args, megatron_args, model): density=1.0, int8_mask=None, ) - return get_megatron_optimizer( + return get_megatron_optimizer_local_sgd( config, model, megatron_args.no_wd_decay_cond, @@ -277,8 +325,6 @@ def prepare_optimizer(train_args, args, megatron_args, model): megatron_args.lr_mult, ) else: - from megatron.core.optimizer import get_megatron_optimizer - return get_megatron_optimizer( config, model, @@ -341,14 +387,18 @@ def model_provider_func(pre_process=True, post_process=True, add_encoder=True, a return model +# override megatron/training/training::should_disable_forward_pre_hook() in Megatron (>=0.12) +def should_disable_forward_pre_hook(): + """Block forward pre-hook for certain configurations.""" + args = get_args() + return not getattr(args, "use_custom_fsdp", False) and args.use_distributed_optimizer and args.overlap_param_gather + + class AtorchMegatronEngine(AtorchTrainEngine): @inject def __init__( self, train_args: AtorchTrainingArgs, - # model: MegatronModule, - # optimizer: MegatronOptimizer, - # scheduler: OptimizerParamScheduler, dataloaders: Union[AtorchMegatronDataloader, tuple, None], train_state: AtorchTrainerState, ckpt_saver: MegatronCkptSaver = Provide[AtorchTrainerContainer.ckpt_saver], @@ -370,7 +420,7 @@ def __init__( self.megatron_args = train_args.megatron_args() - self.initialize(self.megatron_args, self.train_args.use_deterministic_algorithms) + self.initialize() self.num_floating_point_operations_since_last_log_event = 0.0 @@ -418,15 +468,19 @@ def __init__( ) except Exception: raise ValueError("not support swap_attention") + + timers = get_timers() + + # Model, optimizer, and learning rate. + timers("model-and-optimizer-setup", log_level=0).start(barrier=True) ( self.module, self.optimizer, self.scheduler, ) = self.prepare_model_optimizer_scheduler(resume_from_checkpoint=resume_from_checkpoint) + timers("model-and-optimizer-setup").stop() if self.train_args.convert_checkpoint: - from megatron.training.checkpointing import save_checkpoint - args.save = self.train_args.output_dir_for_converting args.async_save = False # asynchronous save is not supported when converting checkpoint @@ -480,11 +534,13 @@ def __init__( tp_group=None, ) + # Moe_routing_map save and load + self.MoELayerDict: Dict[str, MoELayer] = {} + self.saveIter: List[int] = [] + # TODO: define a function to unify barrier operator. torch.distributed.barrier() - # args.iteration = 0 - # args.num_floating_point_operations_so_far = 0 args.model_return_dict = None if self.megatron_args.custom_train_step_class is not None: if self.megatron_args.custom_train_step_kwargs is None: @@ -508,11 +564,13 @@ def __init__( self.eval_total_loss_dict = {} # type: ignore[var-annotated] self.report_memory_flag = True + self.pre_hook_enabled = False self.num_floating_point_operations_so_far = args.num_floating_point_operations_so_far self.module_config = None self.training_log_args = None self.custom_training_log_dict = None self.num_microbatches = get_num_microbatches() + self.skipped_iter = 0 if delay_building_dataloader: self.build_dataloader() @@ -529,100 +587,151 @@ def __init__( torch.cuda.set_rng_state(self.rng_state_from_ckpt["cuda_rng_state"]) tensor_parallel.get_cuda_rng_tracker().set_states(self.rng_state_from_ckpt["rng_tracker_states"]) + # set train_args + self.train_args.per_device_train_batch_size = args.micro_batch_size * get_num_microbatches() + self.train_args.per_device_eval_batch_size = ( + args.global_batch_size // args.data_parallel_size + ) # Don't consider batch size warmup + write_args_to_tensorboard() + # Print setup timing. + log_rank_0("done with setup ...") + timers.log(["model-and-optimizer-setup", "train/valid/test-data-iterators-setup"], barrier=True) + def _validate_train_args( self, ): args = get_args() - self.train_args.per_device_train_batch_size = args.micro_batch_size * get_num_microbatches() - self.train_args.per_device_eval_batch_size = args.micro_batch_size * get_num_microbatches() if args.profile: raise ValueError( "Please set atorch trainer's 'profiler_type' arg, megatron profiler is disabled by" "atorch trainer." ) - # if is_torch_npu_available(): - # logger.warning("nsys profiling is not supported on NPU env.") - # else: - # if self.train_args.profiler_type is None: - # self.train_args.profiler_type = "nsys" - # elif self.train_args.profiler_type != "nsys": - # raise ValueError( - # f"Can't use {self.train_args.profiler_type} and nsys to profile " - # "simultaneously. You can set atorch trainer's 'profiler_type' arg " - # "to 'nsys' or not set Megatron's 'profile' arg." - # ) - - def _assert(condition, arg_name): - assert condition, f"{arg_name} is not supported in AtorchTrainer." - - _assert(not args.enable_one_logger, "enable_one_logger") - _assert(not args.adlr_autoresume, "adlr_autoresume") - _assert(not args.exit_signal_handler, "exit_signal_handler") - _assert(args.exit_duration_in_mins is None, "exit_duration_in_mins") - _assert(args.exit_interval is None, "exit_interval") - _assert(not args.vision_pretraining, "vision_pretraining") + + def _assert(condition, feature, recommend_key_value: dict = None): + err_info = f"{feature} is not supported in AtorchTrainerV2." + if recommend_key_value is not None: + arg_name = list(recommend_key_value.keys())[0] + err_info += f" Please set {arg_name} to {recommend_key_value[arg_name]} ." + assert condition, err_info + + _assert(not args.enable_one_logger, "one_logger", {"enable_one_logger": False}) + _assert(not args.adlr_autoresume, "autoresume on adlr cluster", {"adlr_autoresume": False}) + _assert(not args.exit_signal_handler, "exit_signal_handler", {"exit_signal_handler": False}) + _assert(args.exit_duration_in_mins is None, "exit_duration_in_mins", {"exit_duration_in_mins": None}) + _assert(args.exit_interval is None, "exit_interval", {"exit_interval": None}) + _assert(not args.vision_pretraining, "vision_pretraining", {"vision_pretraining": False}) + _assert( + not getattr(args, "enable_ft_package", False), + "NVIDIA Fault Tolerance", + {"enable_ft_package": False}, + ) + _assert( + getattr(args, "non_persistent_ckpt_type", None) is None, + "non-persistent model checkpoints", + {"non_persistent_ckpt_type": None}, + ) + _assert(not args.log_progress, "log_progress", {"log_progress": False}) + _assert( + not getattr(args, "run_workload_inspector_server", False), + "workload inspector", + {"run_workload_inspector_server": False}, + ) + _assert(not getattr(args, "log_straggler", False), "StragglerDetector", {"log_straggler": False}) + _assert(len(getattr(args, "iterations_to_skip", [])) == 0, "iterations_to_skip", {"iterations_to_skip": []}) + _assert( + not getattr(args, "decrease_batch_size_if_needed", False), + "decrease batch size", + {"decrease_batch_size_if_needed": False}, + ) + _assert( + getattr(args, "train_sync_interval", None) is None, + "Training CPU-GPU synchronization interval", + {"train_sync_interval": None}, + ) + # TODO: to support check_weight_hash_across_dp_replicas_interval + _assert( + getattr(args, "check_weight_hash_across_dp_replicas_interval", None) is None, + "check_weight_hash_across_dp_replicas_interval", + {"check_weight_hash_across_dp_replicas_interval": None}, + ) + assert ( args.ckpt_convert_format is None ), "'ckpt_convert_format' is not supported in AtorchTrainer, please use AtorchTrainer's 'convert_checkpoint' instead." # noqa E501 - @staticmethod - def initialize(megatron_args, use_deterministic_algorithms=False): # todo type type: ignore[override] - initialize_megatron( - extra_args_provider=megatron_args.extra_args_provider, - args_defaults=megatron_args.to_dict(), - ignore_unknown_args=True, - use_deterministic_algorithms=use_deterministic_algorithms, - ) + def initialize(self): + # If megatron >= 0.12.0, call original initialize_megatron() function of Megatron + if is_megatron_version_bigger_than("0.12.0"): + # set up args + parsed_args = setup_args( + extra_args_provider=self.megatron_args.extra_args_provider, + args_defaults=self.megatron_args.to_dict(), + ignore_unknown_args=True, + ) + + if self.train_args.use_deterministic_algorithms: + set_deterministic_algorithms(parsed_args) + + if "get_output_layer_ranks" in inspect.signature(_initialize_distributed).parameters: + # Megatron EA version + initialize_megatron( + extra_args_provider=self.megatron_args.extra_args_provider, + args_defaults=self.megatron_args.to_dict(), + ignore_unknown_args=True, + get_embedding_ranks=self.megatron_args.get_embedding_ranks, + get_position_embedding_ranks=self.megatron_args.get_position_embedding_ranks, + get_output_layer_ranks=self.megatron_args.get_output_layer_ranks, + parsed_args=parsed_args, + ) + else: + initialize_megatron( + extra_args_provider=self.megatron_args.extra_args_provider, + args_defaults=self.megatron_args.to_dict(), + ignore_unknown_args=True, + get_embedding_ranks=self.megatron_args.get_embedding_ranks, + get_position_embedding_ranks=self.megatron_args.get_position_embedding_ranks, + parsed_args=parsed_args, + ) + else: + initialize_megatron_legacy( + extra_args_provider=self.megatron_args.extra_args_provider, + args_defaults=self.megatron_args.to_dict(), + ignore_unknown_args=True, + get_embedding_ranks=self.megatron_args.get_embedding_ranks, + get_position_embedding_ranks=self.megatron_args.get_position_embedding_ranks, + get_output_layer_ranks=self.megatron_args.get_output_layer_ranks, + use_deterministic_algorithms=self.train_args.use_deterministic_algorithms, + ) # Set pytorch JIT layer fusion options and warmup JIT functions. set_jit_fusion_options() - if use_deterministic_algorithms: + if self.train_args.use_deterministic_algorithms: TORCH_MAJOR = int(torch.__version__.split(".")[0]) TORCH_MINOR = int(torch.__version__.split(".")[1]) if (TORCH_MAJOR > 1) or (TORCH_MAJOR == 1 and TORCH_MINOR >= 10): torch._C._jit_set_nvfuser_enabled(False) - def post_training_step(self): - try: - from megatron.training.training import post_training_step_callbacks - - post_training_step_callbacks( - model=self.module, - optimizer=self.optimizer, - opt_param_scheduler=self.scheduler, - iteration=self.iteration, - prof=None, - num_floating_point_operations_since_last_log_event=self.num_floating_point_operations_since_last_log_event, # noqa E501 - ) - - megatron_args = get_args() - if self.iteration % megatron_args.log_interval == 0 and megatron_args.log_straggler: - self.num_floating_point_operations_since_last_log_event = 0.0 - - except ImportError: - pass - def prepare_model_optimizer_scheduler(self, resume_from_checkpoint=None): logger.info("Preparing model optimizer scheduler") args = get_args() - megatron_args = self.megatron_args timers = get_timers() - if megatron_args.custom_prepare_model_function is not None: - if megatron_args.custom_model_provider_function is None: + # TODO: remove some redundant code. + if self.megatron_args.custom_prepare_model_function is not None: + if self.megatron_args.custom_model_provider_function is None: raise ValueError( "You must provide a `custom_model_provider_function` when using a `custom_prepare_model_function`." ) - custom_model_provider_func = megatron_args.custom_model_provider_function - model = megatron_args.custom_prepare_model_function(custom_model_provider_func) + model_provider_func_ = self.megatron_args.custom_model_provider_function + model = self.megatron_args.custom_prepare_model_function(model_provider_func_) else: model_type = ModelType.encoder_or_decoder if args.model_type_name == "t5": model_type = ModelType.encoder_and_decoder - if megatron_args.custom_model_provider_function is not None: - model_provider_func_ = megatron_args.custom_model_provider_function + if self.megatron_args.custom_model_provider_function is not None: + model_provider_func_ = self.megatron_args.custom_model_provider_function else: model_provider_func_ = model_provider_func # NOTE similar to get optimizer, hack here, and need to verify the existance of use_local_sgd @@ -635,10 +744,56 @@ def prepare_model_optimizer_scheduler(self, resume_from_checkpoint=None): model = get_model(model_provider_func_, model_type) - optimizer = prepare_optimizer(self.train_args, args, megatron_args, model) + unwrapped_model = unwrap_model(model) + + optimizer = prepare_optimizer(self.train_args, self.megatron_args, model) scheduler = get_optimizer_param_scheduler(optimizer) - if resume_from_checkpoint is not None: + try: + from ant_utils.dcp_utils.dcp_utils import patch_torch + + patch_torch(patch_find_nd_overlapping_shards=self.train_args.patch_find_nd_overlapping_shards) + except ImportError: + logger.warning( + "Unable to import patch_torch from ant_utils.dcp_utils.dcp_utils, you might not using " + "available megatron version. If you want to use a speedup version of megatron," + "please read atorch doc to use a proper megatron version." + ) + + if is_megatron_version_bigger_than("0.9.0") and args.moe_use_upcycling: + from megatron.core.transformer.moe import upcycling_utils + + torch.distributed.barrier() + assert not checkpoint_exists(args.save), ( + "The upcycling destination directory already exists. " + "Please check if --moe-use-upcycling is mistakenly enabled. " + "Upcycling should only be set for the first run when converting the dense model. " + "All subsequent runs should remove this flag. " + ) + num_experts = args.num_experts + args.num_experts = None + expert_model_parallel_size = args.expert_model_parallel_size + args.expert_model_parallel_size = 1 + dense_model_for_upcycling = get_model(model_provider_func_, model_type) + args.num_experts = num_experts + args.expert_model_parallel_size = expert_model_parallel_size + _, args.num_floating_point_operations_so_far = upcycling_utils.load_and_upcycle_model( + load_checkpoint, + unwrapped_model, + dense_model_for_upcycling, + load_kwargs={"model": dense_model_for_upcycling, "optimizer": None, "opt_param_scheduler": None}, + ) + args.iteration = 1 + save_checkpoint(args.iteration, model, None, None, args.num_floating_point_operations_so_far) + torch.distributed.barrier() + del dense_model_for_upcycling + if (args.fp16 or args.bf16) and optimizer is not None: + optimizer.reload_model_params() + log_rank_0(f"Upcycled checkpoint saved to {args.save}") + + if resume_from_checkpoint is not None and not ( + is_megatron_version_bigger_than("0.9.0") and args.moe_use_upcycling + ): timers("load-checkpoint", log_level=0).start(barrier=True) if isinstance(resume_from_checkpoint, str): @@ -666,20 +821,19 @@ def prepare_model_optimizer_scheduler(self, resume_from_checkpoint=None): args.iteration = 0 args.num_floating_point_operations_so_far = 0 - unwrapped_model = unwrap_model(model) - # get model without FP16 and/or DDP wrappers if ( args.iteration == 0 and len(unwrapped_model) == 1 and hasattr(unwrapped_model[0], "init_state_dict_from_bert") ): - print_rank_0("Initializing ICT from pretrained BERT model") + log_rank_0("Initializing ICT from pretrained BERT model") unwrapped_model[0].init_state_dict_from_bert() if args.fp16: optimizer.reload_model_params() self.iteration = args.iteration + self.start_iteration = args.iteration self.num_floating_point_operations_so_far = args.num_floating_point_operations_so_far args.global_step = args.iteration @@ -698,13 +852,23 @@ def get_dataloader(self, name=None): ) def build_dataloader(self): - def _build_dataloader(): + def _build_dataloader(vp_stage=None): if self.megatron_args.custom_megatron_dataloaders_provider_function is not None: - ( - train_data_iterator, - valid_data_iterator, - test_data_iterator, - ) = self.megatron_args.custom_megatron_dataloaders_provider_function() + if ( + "vp_stage" + in inspect.signature(self.megatron_args.custom_megatron_dataloaders_provider_function).parameters + ): + ( + train_data_iterator, + valid_data_iterator, + test_data_iterator, + ) = self.megatron_args.custom_megatron_dataloaders_provider_function(vp_stage=vp_stage) + else: + ( + train_data_iterator, + valid_data_iterator, + test_data_iterator, + ) = self.megatron_args.custom_megatron_dataloaders_provider_function() return train_data_iterator, valid_data_iterator, test_data_iterator elif self.megatron_args.custom_megatron_datasets_provider_function is not None: ( @@ -722,6 +886,8 @@ def _build_dataloader(): ) = _prepare_megaton_dataloader(self.train_args, self.dataloaders) # self._dataloaders.extend(_prepare_megaton_dataloader(self.train_args, self.dataloaders)) + is_post_training = self.train_args.finetune_type is not None + args = get_args() timers = get_timers() @@ -732,7 +898,14 @@ def _build_dataloader(): test_data_iterator = [] for i in range(args.virtual_pipeline_model_parallel_size): mpu.set_virtual_pipeline_model_parallel_rank(i) - iterators = _build_dataloader() + + if is_post_training: + vp_stage = None + else: + unwrapped_model = unwrap_model(self.module[i]) + vp_stage = unwrapped_model.vp_stage if hasattr(unwrapped_model, "vp_stage") else None + + iterators = _build_dataloader(vp_stage=vp_stage) train_data_iterator.append(iterators[0]) valid_data_iterator.append(iterators[1]) test_data_iterator.append(iterators[2]) @@ -744,8 +917,6 @@ def _build_dataloader(): ) = _build_dataloader() timers("train/valid/test-data-iterators-setup").stop() - is_post_training = self.train_args.finetune_type is not None - if is_post_training: # In post-training scene, we need to calculate how many global steps are required to iterate through the # dataloader. Batch size warmup will increase the complexity of calculating the steps, so we don't consider @@ -782,61 +953,14 @@ def _build_dataloader(): torch.distributed.barrier() - # TODO(@jinshi.cl): Should be called at init, instead when AtorchMegatronEngine.train() or - # AtorchMegatronEngine.eval() - def get_module_config(self): - args = get_args() - config = get_model_config(self.module[0]) - # Setup some training config params - config.grad_scale_func = self.optimizer.scale_loss - config.timers = get_timers() - if isinstance(self.module[0], MegatronDDP) and args.overlap_grad_reduce: - assert config.no_sync_func is None, ( - "When overlap_grad_reduce is True, config.no_sync_func must be None; " - "a custom no_sync_func is not supported when overlapping grad-reduce" - ) - config.no_sync_func = [model_chunk.no_sync for model_chunk in self.module] - if len(self.module) == 1: - config.no_sync_func = config.no_sync_func[0] - should_delay_grad_reduce = ( - args.delay_grad_reduce if hasattr(args, "delay_grad_reduce") else args.align_grad_reduce - ) - if should_delay_grad_reduce: - config.grad_sync_func = [model_chunk.start_grad_sync for model_chunk in self.module] - if len(self.module) == 1: - config.grad_sync_func = config.grad_sync_func[0] - should_delay_param_gather = ( - args.delay_param_gather if hasattr(args, "delay_param_gather") else args.align_param_gather - ) - if args.overlap_param_gather and should_delay_param_gather: - if hasattr(self.optimizer, "finish_param_sync"): - config.param_sync_func = [ - lambda x: self.optimizer.finish_param_sync(model_index, x) - for model_index in range(len(self.module)) - ] - else: - config.param_sync_func = [model_chunk.start_param_sync for model_chunk in self.module] - if len(self.module) == 1: - config.param_sync_func = config.param_sync_func[0] - config.finalize_model_grads_func = finalize_model_grads - return config - def train(self): for model_module in self.module: model_module.train() - if self.module_config is None: - self.module_config = self.get_module_config() - - self.log_eval_results() - def eval(self): for model_module in self.module: model_module.eval() - if self.module_config is None: - self.module_config = self.get_module_config() - def forward(self, data_iterator): # During training, we use train_step() # model(**batch_data) performs following operations by delegating it to `self.train_step`: @@ -858,17 +982,27 @@ def forward(self, data_iterator): args.forward_mode = "train" # Update number of microbatches first without consistency. Then run consistency check # to make sure training configuration is still valid. - update_num_microbatches(args.consumed_train_samples, consistency_check=False) + update_num_microbatches(args.consumed_train_samples, consistency_check=False, verbose=True) if get_num_microbatches() != self.num_microbatches and self.iteration != 0: assert ( get_num_microbatches() > self.num_microbatches ), "number of microbatches should be increasing due to batch size rampup" self.num_microbatches = get_num_microbatches() - update_num_microbatches(args.consumed_train_samples, consistency_check=True) + update_num_microbatches(args.consumed_train_samples, consistency_check=True, verbose=True) args.curr_iteration = self.iteration - loss_dict, skipped_iter, grad_norm, num_zeros_in_grad = self.train_step(data_iterator) + loss_dict, self.skipped_iter, grad_norm, num_zeros_in_grad = self.train_step(data_iterator) self.iteration += 1 + + # save routing map for moe-layer + if getattr(args, "moe_router_save", False): + for router_save_iter in self.saveIter: + if self.iteration == router_save_iter: + for path, layer in self.MoELayerDict.items(): + layer.save_routing_map( + router_save_iter, path, args.moe_router_save_dir, args.moe_splits_save_dir + ) + args.global_step = self.iteration batch_size = mpu.get_data_parallel_world_size() * args.micro_batch_size * get_num_microbatches() args.consumed_train_samples += batch_size @@ -876,7 +1010,11 @@ def forward(self, data_iterator): self.num_floating_point_operations_so_far += num_floating_point_operations_in_batch self.num_floating_point_operations_since_last_log_event += num_floating_point_operations_in_batch - loss_scale = self.optimizer.get_loss_scale().item() + if getattr(self.optimizer, "is_stub_optimizer", False): + loss_scale = 1.0 + else: + loss_scale = self.optimizer.get_loss_scale().item() + params_norm = None if args.log_params_norm: params_norm = calc_params_l2_norm(self.module) @@ -905,11 +1043,11 @@ def forward(self, data_iterator): self.iteration, loss_scale, self.report_memory_flag, - skipped_iter, + self.skipped_iter, grad_norm, params_norm, - params_std, num_zeros_in_grad, + params_std, custom_training_log_dict, ] else: @@ -945,32 +1083,6 @@ def forward(self, data_iterator): return self.train_step_handler.model_output_class(loss=loss, logits=logits) return loss - def get_batch_data_iterator(self, batch_data): - args = get_args() - num_microbatches = get_num_microbatches() - data_chunks = [] - if len(batch_data) > 0: - if num_microbatches > 1: - for i in range(0, num_microbatches): - data_chunks.append( - { - k: v[i * args.micro_batch_size : (i + 1) * args.micro_batch_size] - for k, v in batch_data.items() - } - ) - else: - data_chunks = [batch_data] - - if len(self.module) > 1: - batch_data_iterator = ( - [iter(data_chunks) for _ in range(len(self.module))] - if len(batch_data) > 0 - else [None] * len(self.module) - ) - else: - batch_data_iterator = iter(data_chunks) if len(batch_data) > 0 else None - return batch_data_iterator - def train_step(self, data_iterator): """ Training step for Megatron-LM @@ -982,10 +1094,7 @@ def train_step(self, data_iterator): args = get_args() timers = get_timers() - # Set grad to zero. - for model_chunk in self.module: - model_chunk.zero_grad_buffer() - self.optimizer.zero_grad() + self.optimizer_zero_grad() # Forward pass. forward_backward_func = get_forward_backward_func() @@ -1044,6 +1153,11 @@ def train_step(self, data_iterator): if args.empty_unused_memory_level >= 1: torch.cuda.empty_cache() + # Vision gradients. + if args.vision_pretraining and args.vision_pretraining_type == "dino": + unwrapped_model = unwrap_model(self.module[0]) + unwrapped_model.cancel_gradients_last_layer(args.curr_iteration) + # Update parameters. timers("optimizer", log_level=1).start(barrier=args.barrier_with_L1_time) if spike_loss_ratio == 0.0: @@ -1055,6 +1169,26 @@ def train_step(self, data_iterator): update_successful, grad_norm, num_zeros_in_grad = self.optimizer.step() timers("optimizer").stop() + if is_megatron_version_bigger_than("0.11.0"): + from megatron.training.utils import ( + logical_and_across_model_parallel_group, + reduce_max_stat_across_model_parallel_group, + ) + + # when freezing sub-models we may have a mixture of successful and unsucessful ranks, + # so we must gather across mp ranks + update_successful = logical_and_across_model_parallel_group(update_successful) + # grad_norm and num_zeros_in_grad will be None on ranks without trainable params, + # so we must gather across mp ranks + grad_norm = reduce_max_stat_across_model_parallel_group(grad_norm) + if args.log_num_zeros_in_grad: + num_zeros_in_grad = reduce_max_stat_across_model_parallel_group(num_zeros_in_grad) + + # Vision momentum. + if args.vision_pretraining and args.vision_pretraining_type == "dino": + unwrapped_model = unwrap_model(self.module[0]) + unwrapped_model.update_momentum(args.curr_iteration) + # Update learning rate. if update_successful: increment = get_num_microbatches() * args.micro_batch_size * args.data_parallel_size @@ -1079,6 +1213,10 @@ def eval_step(self, data_iterator): args = get_args() + # make validation batch size independent from training batch size + eval_batch_size = args.global_batch_size + eval_num_microbatches = eval_batch_size // (args.micro_batch_size * args.data_parallel_size) + forward_backward_func = get_forward_backward_func() # Don't care about timing during evaluation self.module_config.timers = None @@ -1087,9 +1225,10 @@ def eval_step(self, data_iterator): forward_step_func=self.train_step_handler.get_forward_step_func(), data_iterator=data_iterator, model=self.module, - num_microbatches=get_num_microbatches(), + num_microbatches=eval_num_microbatches, seq_length=args.seq_length, micro_batch_size=args.micro_batch_size, + decoder_seq_length=args.decoder_seq_length, forward_only=True, ) self.module_config.timers = get_timers() @@ -1118,14 +1257,11 @@ def eval_step(self, data_iterator): ) loss_to_log = self.train_step_handler.validation_loss_postprocessing(loss_dicts) - args.consumed_valid_samples += ( - mpu.get_data_parallel_world_size() * args.micro_batch_size * get_num_microbatches() - ) + args.consumed_valid_samples += eval_batch_size return loss_to_log def log_eval_results(self): - args = get_args() if self.iteration == 0 or len(self.eval_total_loss_dict) == 0: return args = get_args() @@ -1166,6 +1302,26 @@ def save_checkpoint( best_model_checkpoint=None, **kwargs, ): + def _save_trainer_state(): + if torch.distributed.get_rank() == 0: + try: + checkpoint_dir_path = self.get_checkpoint_iteration_path_dir(None, return_base_dir=True) + trainer_state_path = str(checkpoint_dir_path.joinpath(TRAINER_STATE_NAME)) + self.train_state.save_to_json(trainer_state_path) + logger.info(f"Successfully save trainer_state.json at {trainer_state_path}") + except Exception as e: + logger.error(f"Fail to save {trainer_state_path}! {e}") + + custom_async_finalize_fn = None + if "custom_async_finalize_fn" in inspect.signature(save_checkpoint).parameters: + custom_async_finalize_fn = _save_trainer_state + kwargs.update(custom_async_finalize_fn=custom_async_finalize_fn) + + timers = get_timers() + timers("save-checkpoint", log_level=0).start(barrier=True) + + if should_disable_forward_pre_hook(): + self.disable_forward_pre_hook() self.ckpt_saver.save( iteration=self.iteration, output_dir=str(output_dir), @@ -1175,12 +1331,19 @@ def save_checkpoint( scheduler=self.scheduler, num_floating_point_operations_so_far=self.num_floating_point_operations_so_far, best_model_checkpoint=best_model_checkpoint, + train_state=self.train_state, + **kwargs, ) + if should_disable_forward_pre_hook(): + self.enable_forward_pre_hook() - if self.train_args.is_main_process: - checkpoint_dir_path = self.get_checkpoint_iteration_path_dir(None, return_base_dir=True) + timers("save-checkpoint").stop(barrier=True) + timers.log(["save-checkpoint"]) - self.train_state.save_to_json(str(checkpoint_dir_path.joinpath(TRAINER_STATE_NAME))) + gc.collect() + + if custom_async_finalize_fn is None: + _save_trainer_state() def get_checkpoint_iteration_path_dir(self, output_dir: Optional[Path], **kwargs) -> Path: """ @@ -1212,7 +1375,11 @@ def load_checkpoint(self, resume_from_ckpt: Optional[Path], model, optimizer=Non args.load = str(resume_from_ckpt) iteration, num_floating_point_operations_so_far = self.ckpt_loader.load( - model=model, optimizer=optimizer, scheduler=scheduler, train_args=self.train_args + model=model, + optimizer=optimizer, + scheduler=scheduler, + train_args=self.train_args, + **kwargs, ) self.iteration = iteration @@ -1223,18 +1390,326 @@ def load_checkpoint(self, resume_from_ckpt: Optional[Path], model, optimizer=Non return iteration, num_floating_point_operations_so_far - def optimizer_step(self): - # Megatron's train_step() contains optimizer.step() + def optimizer_zero_grad(self): + # Set grad to zero. + for model_chunk in self.module: + model_chunk.zero_grad_buffer() + self.optimizer.zero_grad() + + def disable_forward_pre_hook(self, param_sync=True, pre_hook_enabled=None): + if is_megatron_version_bigger_than("0.10.0"): + from megatron.training.training import disable_forward_pre_hook + + if "param_sync" in inspect.signature(disable_forward_pre_hook).parameters: + disable_forward_pre_hook(self.module, param_sync=param_sync) + else: + disable_forward_pre_hook(self.module) + if pre_hook_enabled is not None: + self.pre_hook_enabled = pre_hook_enabled + else: + self.optimizer.disable_pre_hook() + + def enable_forward_pre_hook(self, pre_hook_enabled=None): + if is_megatron_version_bigger_than("0.10.0"): + from megatron.training.training import enable_forward_pre_hook + + enable_forward_pre_hook(self.module) + if pre_hook_enabled is not None: + self.pre_hook_enabled = pre_hook_enabled + else: + self.optimizer.enable_pre_hook() + + +class MegatronCallback(AtorchTrainerCallback): + """ + A [`AtorchTrainerCallback`] that supplements Megatron training process. + """ + + def __init__(self, train_engine: AtorchMegatronEngine): + super().__init__() + self.train_engine = train_engine + self.eval_type = None + + def on_train_begin( + self, + atorch_training_args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + """ + Event called at the beginning of training, after Megatron initialization and building dataloader. + """ + module = self.train_engine.module + optimizer = self.train_engine.optimizer + + args = get_args() + config = get_model_config(module[0]) + # Setup some training config params + config.grad_scale_func = optimizer.scale_loss + config.timers = get_timers() + if isinstance(module[0], MegatronDDP) and args.overlap_grad_reduce: + assert config.no_sync_func is None, ( + "When overlap_grad_reduce is True, config.no_sync_func must be None; " + "a custom no_sync_func is not supported when overlapping grad-reduce" + ) + config.no_sync_func = [model_chunk.no_sync for model_chunk in module] + if len(module) == 1: + config.no_sync_func = config.no_sync_func[0] + should_delay_grad_reduce = ( + args.delay_grad_reduce if hasattr(args, "delay_grad_reduce") else args.align_grad_reduce + ) + if should_delay_grad_reduce: + config.grad_sync_func = [model_chunk.start_grad_sync for model_chunk in module] + if len(module) == 1: + config.grad_sync_func = config.grad_sync_func[0] + should_delay_param_gather = ( + args.delay_param_gather if hasattr(args, "delay_param_gather") else args.align_param_gather + ) + if args.overlap_param_gather and should_delay_param_gather: + if hasattr(optimizer, "finish_param_sync"): + config.param_sync_func = [ + lambda x: optimizer.finish_param_sync(model_index, x) for model_index in range(len(module)) + ] + else: + config.param_sync_func = [model_chunk.start_param_sync for model_chunk in module] + if len(module) == 1: + config.param_sync_func = config.param_sync_func[0] + config.finalize_model_grads_func = finalize_model_grads + + # Disable forward pre-hook to start training to ensure that errors in checkpoint loading + # or random initialization don't propagate to all ranks in first all-gather (which is a + # no-op if things work correctly). + if should_disable_forward_pre_hook(): + """Block forward pre-hook for certain configurations.""" + self.train_engine.disable_forward_pre_hook(param_sync=False, pre_hook_enabled=False) + # Also remove param_sync_func temporarily so that sync calls made in + # `forward_backward_func` are no-ops. + self.train_engine.param_sync_func = config.param_sync_func + config.param_sync_func = None + + # print model config + if atorch_training_args.is_local_main_process and atorch_training_args.debug_switch.get("print_config", True): + logger.info(f"************************ {config.__class__.__name__} ************************") + logger.info(f"model config: {config}") + logger.info(f"************************ {config.__class__.__name__} ************************") + + self.train_engine.module_config = config + + # ============================================= + # moe_routing_map load and save + # ============================================= + # Save moelayer routing map + if getattr(args, "moe_router_save", False): + # if save routing map dir and iters must be specified + if ( + args.moe_router_save_dir is None + or args.moe_router_save_iters is None + or args.moe_splits_save_dir is None + ): + raise ValueError( + "When --moe-router-save is True, " + "--moe-router-save-dir, --moe-router-save-iters " + "and --moe-splits-save_dir must be specified" + ) + if args.moe_token_dispatcher_type != "alltoall": + raise ValueError("--moe-router-save only supports --moe-token-dispatcher-type=alltoall") + # mkdir if not exists + os.makedirs(args.moe_router_save_dir, exist_ok=True) + os.makedirs(args.moe_splits_save_dir, exist_ok=True) + # save iters + str_list = args.moe_router_save_iters.split(",") + int_list = [int(x) for x in str_list] + self.train_engine.saveIter = int_list + + # TODO: parse related args and check + # load moelayer routing map + if getattr(args, "moe_router_load", False): + if args.moe_router_load_dir is None: + raise ValueError("When --moe-router-load is True, --moe-router-load-dir must be specified") + # find all the moe-layers + if getattr(args, "moe_router_save", False) or getattr(args, "moe_router_load", False): + MoELayerDict = {} + + def find_layers_with_path(module, target_class=MoELayer, prefix=""): + # check if the module is target + if isinstance(module, target_class): + yield (prefix, module) + + # sub_modules + for name, child in module.named_children(): + current_path = f"{prefix}.{name}" if prefix else name + yield from find_layers_with_path(child, target_class, current_path) + + for path, layer in find_layers_with_path(module[0]): + key = f"{path.replace('.', '_')}_layerid_{layer.layer_number}" + MoELayerDict[key] = layer + self.train_engine.MoELayerDict = MoELayerDict + + # load moe-routing-map + if getattr(args, "moe_router_load", False): + + def parse_routing_map_filename(file_path): + """parse file name""" + dir_name = os.path.basename(os.path.dirname(file_path)) # iter_100 + file_name = os.path.basename(file_path) # layer_encoder_layer_id_2_dp_rank_0_routing_map.pt + # parse iteration + try: + iteration = int(dir_name.replace("iter_", "")) + except ValueError: + raise ValueError(f"Invalid iteration directory name: {dir_name}") + # parse other information + pattern = ( + r"layer_(?P\w+)_layer_id_(?P\d+)_dp_rank_(?P\d+)_routing_map\.pt" + ) + match = re.search(pattern, file_name) + if not match: + logger.error( + f"Invalid filename format: '{file_name}'\n" + "Expected format: 'layer__layer_id__dp_rank__routing_map.pt'" + ) + raise ValueError(f"Invalid filename format: {file_name}") + return { + "layer_name": match.group("layer_name"), + "layer_id": int(match.group("layer_id")), + "dp_rank": int(match.group("dp_rank")), + "iteration": iteration, + "file_path": file_path, + } + + def process_iteration_directory(iter_dir, layers_dict): + if not os.path.isdir(iter_dir): + raise ValueError(f"Directory does not exist: {iter_dir}") + # acquire all .pt files + pt_files = [os.path.join(iter_dir, f) for f in os.listdir(iter_dir) if f.endswith(".pt")] + for file_path in pt_files: + params = parse_routing_map_filename(file_path) + layer = MoELayerDict.get(f'{params["layer_name"]}') + if layer is None: + continue + layer.token_dispatcher.load_routing_map(file_path, params["dp_rank"]) + return + + # load ckpts + process_iteration_directory(args.moe_router_load_dir, MoELayerDict) + for layer in MoELayerDict.values(): + assert layer.token_dispatcher.preload_routing_map is not None, "Routing map not loaded" + + ############################################### + + def on_step_end( + self, + atorch_training_args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + """ + Event called at the end of a training step, after forward+backward+optimizer.step, + but before logging/evaluate/save_checkpoint + """ + args = get_args() + # Enable forward pre-hooks after first set of forward and backward passes. + # When running in fp16, skip all NaN iterations until steady-state loss scaling value + # is reached. + if args.curr_iteration == self.train_engine.start_iteration: + if self.train_engine.skipped_iter: + # Only enable forward pre-hook after a training step has successfully run. Relevant + # for fp16 codepath where first XX iterations are skipped until steady-state loss + # scale value is reached. + self.train_engine.start_iteration = args.curr_iteration + 1 + else: + # Enable forward pre-hook after training step has successfully run. All subsequent + # forward passes will use the forward pre-hook / `param_sync_func` in + # `forward_backward_func`. + if should_disable_forward_pre_hook(): + self.train_engine.enable_forward_pre_hook(pre_hook_enabled=True) + self.train_engine.module_config.param_sync_func = self.train_engine.param_sync_func # type: ignore[attr-defined] # noqa E501 + + def on_evaluate_begin( # type: ignore[override] + self, + atorch_training_args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + """ + Event called at the beginning of evaluation. + """ + # "eval" or "test" + self.eval_type = kwargs["eval_type"] if "eval_type" in kwargs and kwargs["eval_type"] is not None else "eval" + + if should_disable_forward_pre_hook(): + self.train_engine.disable_forward_pre_hook(pre_hook_enabled=False) + + args = get_args() + timers = get_timers() + timers(f"{self.eval_type}-time", log_level=0).start(barrier=True) + + if args.vision_pretraining and args.vision_pretraining_type == "dino": + from megatron.legacy.model.vision.knn_monitor import compute_feature_bank + + compute_feature_bank(self.train_engine.module) + + def on_evaluate( + self, + atorch_training_args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + """ + Event called after evaluation. + """ + # Print eval result + self.train_engine.log_eval_results() + + timers = get_timers() + timers(f"{self.eval_type}-time").stop() + timers.log([f"{self.eval_type}-time"]) + + if should_disable_forward_pre_hook(): + self.train_engine.enable_forward_pre_hook(pre_hook_enabled=True) + + def on_save_begin( # type: ignore[override] + self, + atorch_training_args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + """ + Event called at the beginning of saving. + """ pass - # return self.optimizer.step() - def scheduler_step(self): - # Megatron's train_step() contains scheduler.step() + def on_save( + self, + atorch_training_args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + """ + Event called after a checkpoint save. + """ pass - # return self.scheduler.step() - def optimizer_zero_grad(self): - return self.optimizer.zero_grad() + def on_train_end( + self, + atorch_training_args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + """ + Event called at the end of training. + """ + # Flush TensorBoard, WandB writers and one-logger. + writer = get_tensorboard_writer() + if writer: + writer.flush() - def backward(self, loss): - pass + # Close out pre-hooks if using distributed optimizer and overlapped param gather. + if self.train_engine.pre_hook_enabled: + self.train_engine.disable_forward_pre_hook() diff --git a/atorch/trainer/trainer_callback.py b/atorch/trainer/trainer_callback.py index 219d289..923b969 100644 --- a/atorch/trainer/trainer_callback.py +++ b/atorch/trainer/trainer_callback.py @@ -1,6 +1,9 @@ import json +import os from dataclasses import dataclass, fields +from typing import List +import torch from tqdm import tqdm from transformers.trainer_callback import CallbackHandler, TrainerCallback, TrainerControl, TrainerState from transformers.trainer_utils import IntervalStrategy @@ -8,9 +11,11 @@ from atorch.common.log_utils import default_logger as logger from atorch.trainer.args import AtorchTrainingArgs from atorch.trainer.atorch_args import AtorchArguments +from atorch.trainer.atorch_profiler import get_profiler from atorch.trainer.utils import DistributedType from atorch.trainer.utils import IntervalStrategy as AtorchIntervalStrategy -from atorch.utils.import_util import is_megatron_lm_available +from atorch.utils.dynamic_profiler import FileWatcher +from atorch.utils.import_util import is_megatron_lm_available, is_torch_npu_available if is_megatron_lm_available(): try: @@ -19,6 +24,13 @@ from megatron.core.num_microbatches_calculator import get_current_global_batch_size +if is_torch_npu_available(): + try: + from torch_npu.profiler import dynamic_profile as dp + except ImportError: + dp = None + + @dataclass class AtorchTrainerState(TrainerState): steps_in_epoch: int = 0 @@ -125,10 +137,18 @@ def call_event_safely(self, event, args, state, control, **kwargs): if not hasattr(callback, event): continue + # 'tokenizer' is removed from CallbackHandler after transformers 4.45.0 + if hasattr(self, "tokenizer"): + kwargs.update(tokenizer=self.tokenizer) result = getattr(callback, event)( args, state, control, + model=self.model, + optimizer=self.optimizer, + lr_scheduler=self.lr_scheduler, + train_dataloader=self.train_dataloader, + eval_dataloader=self.eval_dataloader, **kwargs, ) # A Callback can skip the return of `control` if it doesn't change it. @@ -198,6 +218,32 @@ class FlowCallbackV2(AtorchTrainerCallback): A [`AtorchTrainerCallback`] that handles the default flow of the training loop for logs, evaluation and checkpoints. """ + def __init__(self): + super().__init__() + self.save_ckpt_file_monitor = None + self.save_at_dynamic_steps = [] + + def on_train_begin( + self, args: AtorchTrainingArgs, state: AtorchTrainerState, control: AtorchTrainerControl, **kwargs + ): + if args.dynamic_save_config_path is not None: + self.save_ckpt_file_monitor = FileWatcher(args.dynamic_save_config_path, expire_time=0) + if os.path.exists(args.dynamic_save_config_path): + self.save_ckpt_file_monitor._last_mtime = os.path.getmtime(args.dynamic_save_config_path) + with open(args.dynamic_save_config_path, "r") as f: + dynamic_saving_config = json.loads(f.read()) + if "save_at_dynamic_steps" in dynamic_saving_config and isinstance( + dynamic_saving_config["save_at_dynamic_steps"], List + ): + self.save_at_dynamic_steps = dynamic_saving_config["save_at_dynamic_steps"] + logger.info( + f"[Dynamic saving] checkpoint at global steps {self.save_at_dynamic_steps} will be saved!" + ) + else: + logger.info( + f"[WARNING] You should set 'save_at_dynamic_steps' to a list in {args.dynamic_save_config_path}" # noqa: E501 + ) + def on_step_begin( self, args: AtorchTrainingArgs, state: AtorchTrainerState, control: AtorchTrainerControl, **kwargs ): @@ -242,6 +288,23 @@ def _judge_save_ckpt_by_samples(): should_save = True return should_save + def _save_at_dynamic_steps(): + if self.save_ckpt_file_monitor is None: + return False + file_content, _ = self.save_ckpt_file_monitor.read_if_modified() + if file_content is not None: + config_dict = json.loads(file_content) + if "save_at_dynamic_steps" in config_dict and isinstance(config_dict["save_at_dynamic_steps"], list): + self.save_at_dynamic_steps = config_dict["save_at_dynamic_steps"] + logger.info( + f"[Dynamic saving] checkpoint at global steps {self.save_at_dynamic_steps} will be saved!" + ) + else: + logger.info( + f"[WARNING] You should set 'save_at_dynamic_steps' to a list in {args.dynamic_save_config_path}" + ) + return state.global_step in self.save_at_dynamic_steps + # Save if ( ( @@ -259,6 +322,7 @@ def _judge_save_ckpt_by_samples(): args.extra_save_frequency_in_epoch is not None and state.current_step_in_epoch in args.extra_save_frequency_in_epoch ) + or _save_at_dynamic_steps() ): control.should_save = True # Extra judge about test. @@ -341,3 +405,64 @@ def on_predict(self, args: AtorchTrainingArgs, state: AtorchTrainerState, contro if self.prediction_bar is not None: self.prediction_bar.close() self.prediction_bar = None + + +class ProfilerCallback(AtorchTrainerCallback): + """ + A [`AtorchTrainerCallback`] that execute profiling. + """ + + def on_train_begin( + self, + args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + self.prof = get_profiler(args) + if hasattr(self.prof, "start"): + self.prof.start() + + def on_train_end( + self, + args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + if hasattr(self.prof, "stop"): + self.prof.stop() + + def on_step_begin( + self, + args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + if ( + args.profiler_type == "nsys" + and state.global_step == args.profile_step_start + and args.process_index in args.profile_ranks + ): + torch.cuda.cudart().cudaProfilerStart() + torch.autograd.profiler.emit_nvtx(record_shapes=True).__enter__() + + def on_step_end( + self, + args: AtorchTrainingArgs, + state: AtorchTrainerState, + control: AtorchTrainerControl, + **kwargs, + ): + if ( + args.profiler_type == "nsys" + and state.global_step == args.profile_step_end + and args.process_index in args.profile_ranks + ): + torch.cuda.cudart().cudaProfilerStop() + elif args.profiler_type == "hw_dp": + if dp is not None: + dp.step() + elif self.prof is not None and hasattr(self.prof, "step"): + self.prof.step() diff --git a/atorch/trainer/utils.py b/atorch/trainer/utils.py index d63df1a..a6358d7 100644 --- a/atorch/trainer/utils.py +++ b/atorch/trainer/utils.py @@ -275,6 +275,7 @@ def write_dict_to_tensorboard(writer, wandb_writer, index, metrics, prefix=""): wandb_writer.log({key_to_write: value}, index) +# TODO:(L1) Use original training_log on Megatron>=0.12 def training_log( loss_dict, total_loss_dict, @@ -286,8 +287,8 @@ def training_log( skipped_iter, grad_norm, params_norm, - params_std, num_zeros_in_grad, + params_std, custom_metrics, ): """Log training information such as losses, timing, ....""" @@ -367,6 +368,11 @@ def training_log( total_iterations = total_loss_dict[advanced_iters_key] + total_loss_dict[skipped_iters_key] + if is_megatron_version_bigger_than("0.11.0"): + from megatron.training.utils import reduce_max_stat_across_model_parallel_group + + # learning rate will be None on ranks without trainable params, so we must gather across mp ranks + learning_rate = reduce_max_stat_across_model_parallel_group(learning_rate) # Tensorboard values. # Timer requires all the ranks to call. if args.log_timers_to_tensorboard and (iteration % args.tensorboard_log_interval == 0): @@ -387,24 +393,22 @@ def training_log( if wandb_writer: wandb_writer.log({"samples vs steps": args.consumed_train_samples}, iteration) - should_log_learning_rate_to_tensorboard = ( - args.log_learning_rate_to_tensorboard if hasattr(args, "log_learning_rate_to_tensorboard") else True - ) - if should_log_learning_rate_to_tensorboard: - writer.add_scalar("learning-rate", learning_rate, iteration) - if args.decoupled_lr is not None: - writer.add_scalar("decoupled-learning-rate", decoupled_learning_rate, iteration) - writer.add_scalar("learning-rate vs samples", learning_rate, args.consumed_train_samples) - if wandb_writer: - wandb_writer.log({"learning-rate": learning_rate}, iteration) - should_log_batch_size_to_tensorboard = ( - args.log_batch_size_to_tensorboard if hasattr(args, "log_batch_size_to_tensorboard") else True - ) - if should_log_batch_size_to_tensorboard: - writer.add_scalar("batch-size", batch_size, iteration) - writer.add_scalar("batch-size vs samples", batch_size, args.consumed_train_samples) + writer.add_scalar("learning-rate", learning_rate, iteration) + writer.add_scalar("learning-rate vs samples", learning_rate, args.consumed_train_samples) + if wandb_writer: + wandb_writer.log({"learning-rate": learning_rate}, iteration) + if args.decoupled_lr is not None: + writer.add_scalar("decoupled-learning-rate", decoupled_learning_rate, iteration) + if args.skipped_train_samples > 0: + writer.add_scalar("skipped-train-samples", args.skipped_train_samples, iteration) if wandb_writer: - wandb_writer.log({"batch-size": batch_size}, iteration) + wandb_writer.log({"skipped-train-samples": args.skipped_train_samples}, iteration) + + writer.add_scalar("batch-size", batch_size, iteration) + writer.add_scalar("batch-size vs samples", batch_size, args.consumed_train_samples) + if wandb_writer: + wandb_writer.log({"batch-size": batch_size}, iteration) + for key in loss_dict: writer.add_scalar(key, loss_dict[key], iteration) writer.add_scalar(key + " vs samples", loss_dict[key], args.consumed_train_samples) @@ -452,6 +456,11 @@ def training_log( mem_stats["allocated_bytes.all.current"], iteration, ) + writer.add_scalar( + "mem-max-allocated-bytes", + mem_stats["allocated_bytes.all.peak"], + iteration, + ) writer.add_scalar( "mem-allocated-count", mem_stats["allocation.all.current"], @@ -465,9 +474,31 @@ def training_log( if args.num_experts is not None: moe_loss_scale = 1 / get_num_microbatches() - track_moe_metrics(moe_loss_scale, iteration, writer, wandb_writer, total_loss_dict, args.moe_per_layer_logging) - if hasattr(args, "mtp_num_layers") and args.mtp_num_layers is not None: + if is_megatron_version_bigger_than("0.12.0"): + track_names = [] + if args.moe_router_load_balancing_type in ["aux_loss", "seq_aux_loss"]: + track_names.append("load_balancing_loss") + if args.moe_z_loss_coeff is not None: + track_names.append("z_loss") + track_moe_metrics( + loss_scale=moe_loss_scale, + iteration=iteration, + writer=writer, + wandb_writer=wandb_writer, + total_loss_dict=total_loss_dict, + per_layer_logging=args.moe_per_layer_logging, + force_initialize=True, + track_names=track_names, + num_layers=args.num_layers, + moe_layer_freq=args.moe_layer_freq, + ) + else: + track_moe_metrics( + moe_loss_scale, iteration, writer, wandb_writer, total_loss_dict, args.moe_per_layer_logging + ) + + if getattr(args, "mtp_num_layers", None) is not None: from megatron.core.transformer.multi_token_prediction import MTPLossLoggingHelper mtp_loss_scale = 1 / get_num_microbatches() @@ -557,13 +588,14 @@ def training_log( report_memory("(after {} iterations)".format(iteration)) report_memory_flag = False - if hasattr(args, "use_local_sgd") and args.use_local_sgd: - # Lazy import patches for local sgd logging - from atorch.local_sgd.megatron.parallel_state import get_non_data_parallel_group + if not args.log_timers_to_tensorboard or iteration % args.tensorboard_log_interval != 0: + if hasattr(args, "use_local_sgd") and args.use_local_sgd: + # Lazy import patches for local sgd logging + from atorch.local_sgd.megatron.parallel_state import get_non_data_parallel_group - timers.log(timers_to_log, normalizer=args.log_interval, process_group=get_non_data_parallel_group()) - else: - timers.log(timers_to_log, normalizer=args.log_interval) + timers.log(timers_to_log, normalizer=args.log_interval, process_group=get_non_data_parallel_group()) + else: + timers.log(timers_to_log, normalizer=args.log_interval) return report_memory_flag, all_logging_metrics @@ -704,6 +736,8 @@ def scale_main_grad_for_spike_loss( def get_grads_in_optimizer(optimizer: "MegatronOptimizer"): + from atorch.utils.virtual_optimizer.megatron_virtual_optimizer import MegatronVirtualOptimizer + if isinstance(optimizer, (DistributedOptimizer, Float16OptimizerWithFloat16Params)): optimizer._copy_model_grads_to_main_grads() return optimizer.get_main_grads_for_grad_norm() @@ -712,6 +746,8 @@ def get_grads_in_optimizer(optimizer: "MegatronOptimizer"): for single_optimizer in optimizer.chained_optimizers: grads_for_scaling += get_grads_in_optimizer(single_optimizer) return grads_for_scaling + elif isinstance(optimizer, MegatronVirtualOptimizer): + return [] else: raise ValueError(f"Unsupported optimizer type {type(optimizer)}") @@ -724,6 +760,8 @@ def _get_grad_state_model_parallel_group(optimizer: "MegatronOptimizer"): def get_grad_norm_in_optimizer(optimizer: "MegatronOptimizer"): + from atorch.utils.virtual_optimizer.megatron_virtual_optimizer import MegatronVirtualOptimizer + if isinstance(optimizer, (DistributedOptimizer, Float16OptimizerWithFloat16Params)): grads_for_norm = get_grads_in_optimizer(optimizer) return get_grad_norm_fp32( @@ -735,6 +773,8 @@ def get_grad_norm_in_optimizer(optimizer: "MegatronOptimizer"): for single_optimizer in optimizer.chained_optimizers: grad_norms.append(get_grad_norm_in_optimizer(single_optimizer)) return math.sqrt(sum([x**2 for x in grad_norms])) + elif isinstance(optimizer, MegatronVirtualOptimizer): + return 0 else: raise ValueError(f"Unsupported optimizer type {type(optimizer)}") diff --git a/atorch/utils/dynamic_profiler/__init__.py b/atorch/utils/dynamic_profiler/__init__.py index 3afb53b..9915749 100644 --- a/atorch/utils/dynamic_profiler/__init__.py +++ b/atorch/utils/dynamic_profiler/__init__.py @@ -1,3 +1,4 @@ from ._dynamic_profile import init +from ._file_monitor import FileWatcher -__all__ = ["init"] +__all__ = ["init", "FileWatcher"] diff --git a/atorch/utils/dynamic_profiler/_dynamic_profile.py b/atorch/utils/dynamic_profiler/_dynamic_profile.py index db77b70..8fdf54b 100644 --- a/atorch/utils/dynamic_profiler/_dynamic_profile.py +++ b/atorch/utils/dynamic_profiler/_dynamic_profile.py @@ -2,31 +2,80 @@ import functools import json import os +import pickle import socket -from dataclasses import dataclass, field -from typing import Optional +import threading +import time +from dataclasses import asdict, dataclass, field +from datetime import datetime +from enum import Enum +from typing import Dict, List, Optional import torch from torch import profiler +from atorch import local_rank, rank, world_size from atorch.common.log_utils import default_logger as logger from atorch.common.singleton import SingletonMeta -from atorch.distributed.distributed import rank +from atorch.utils.import_util import is_megatron_lm_available -from ._file_monitor import ThreadFileConfigMonitor +from ._file_monitor import ThreadFileConfigMonitor, datetime_field __all__ = ["init"] -def active_kineto() -> bool: - dynolog_flag = os.getenv("KINETO_USE_DAEMON", 0) - try: - dynolog_flag = int(dynolog_flag) - except ValueError: - logger.error("Environment variable KINETO_USE_DAEMON value not valid, will be set to 0 !") - dynolog_flag = 0 +def local_world_size(): + return int(os.getenv("LOCAL_WORLD_SIZE", 1)) + + +@dataclass +class MegatronParallelConfig: + tensor_model_parallel_size: int = 1 + pipeline_model_parallel_size: int = 1 + context_parallel_size: int = 1 + expert_model_parallel_size: int = 1 + expert_tensor_parallel_size: Optional[int] = None + use_tp_pp_dp_mapping: bool = False + encoder_tensor_model_parallel_size: int = 0 + encoder_pipeline_model_parallel_size: int = 0 + world_size: int = 1 + + @classmethod + def from_megatron_args(cls): + if not is_megatron_lm_available(): + logger.warning("MegatronLM is not available, skip getting megatron parallel config") + return None + + try: + from megatron.training.global_vars import get_args # type: ignore[attr-defined] + except ImportError: + logger.warning("Failed to import megatron.training.global_vars, " "skip getting megatron parallel config") + return None + + args = get_args() + + return cls( + tensor_model_parallel_size=args.tensor_model_parallel_size, + pipeline_model_parallel_size=args.pipeline_model_parallel_size, + context_parallel_size=args.context_parallel_size, + expert_model_parallel_size=args.expert_model_parallel_size, + expert_tensor_parallel_size=args.expert_tensor_parallel_size, + use_tp_pp_dp_mapping=args.use_tp_pp_dp_mapping, + encoder_tensor_model_parallel_size=args.encoder_tensor_model_parallel_size, + encoder_pipeline_model_parallel_size=args.encoder_pipeline_model_parallel_size, + world_size=world_size(), + ) + + def to_dict(self) -> dict: + return asdict(self) + + @classmethod + def from_dict(cls, data: dict) -> "MegatronParallelConfig": + return cls(**data) - return dynolog_flag == 1 + +def active_kineto() -> bool: + return os.getenv("KINETO_USE_DAEMON", None) is not None if not active_kineto(): @@ -53,29 +102,55 @@ def wrapper(*args, **kwargs): @dataclass(frozen=True) -class ProfilerConfig: +class OnDemandProfilerConfig: """Immutable profiler configuration.""" - enabled: bool = False + # output config output_dir: str = "" + use_gzip: bool = True + + # mode in ["trace", "dump"] + mode: str = "trace" + + # use start step or start time, if both are set, config will be ignored start_step: int = 0 + start_time: Optional[datetime] = datetime_field(date_format="%Y%m%d%H%M") + + # schedule config schedule_wait: int = 0 schedule_warmup: int = 0 schedule_active: int = 0 - schedule_repeat: int = 1 - schedule_skip_first: int = 0 + + # profiler config with_stack: bool = False with_flops: bool = False with_modules: bool = False record_shapes: bool = False profile_memory: bool = False - # acc_events: bool = False # TODO: add acc_events for new version activities: list = field(default_factory=list) meta_data: dict = field(default_factory=dict) profile_ranks: list = field(default_factory=list) - use_gzip: bool = True + + # dump snapshot config + enabled: str = "all" + context: str = "all" + stacks: str = "all" + max_entries: int = 100000 + # # torch old version.. + # device: "Device" = None + # record_context_cpp: bool = False + # clear_history: bool = False + # compile_context: bool = False + # global_record_annotations: bool = False + + # x_config_update_time will be injected by file monitor + x_config_update_time: datetime = datetime_field(date_format="%Y%m%d%H%M") + # profile session id, if not specified by config file, it will be generated by + # config update time with format %Y%m%d%H%M%S + session_id: str = "" def __post_init__(self): + # convert activities to list of ProfilerActivity activities = self.activities new_activities = [] for activity in activities: @@ -86,25 +161,375 @@ def __post_init__(self): new_activities.append(prof_activity) object.__setattr__(self, "activities", new_activities) + # generate session id if not specified + if self.session_id == "": + object.__setattr__( + self, + "session_id", + self.x_config_update_time.strftime("%Y%m%d%H%M%S"), + ) + def is_valid(self) -> bool: """Check if the configuration is valid.""" + if self.start_step > 0 and self.start_time is not None: + logger.warning("start_step and start_time can not be set at the same time") + return False + if self.start_step == 0 and self.start_time is None: + logger.warning("start_step and start_time can not be set to 0") + return False + + if self.profile_memory: + if not self.with_stack or not self.record_shapes: + logger.warning("Profile memory is enabled, but with_stack or record_shapes is not enabled") + return False + + # Check mode in ["trace", "dump"] + if self.mode not in ["trace", "dump"]: + logger.warning(f"Profile mode `{self.mode}` only support `trace` or `dump`") + return False + return ( - self.enabled - and self.schedule_active > 0 + self.schedule_active > 0 and self.output_dir != "" and len(self.activities) > 0 and len(self.profile_ranks) > 0 ) +class LocalShmStore: + def __init__(self, store_path: str) -> None: + self.store_path = store_path + self.base_dir = os.path.join("/dev/shm", f"{store_path}") + + def set_value(self, key: str, value: str): + path = os.path.join(self.base_dir, f"{key}") + os.makedirs(os.path.dirname(path), exist_ok=True) + with open(path, "w") as f: + f.write(value) + + def set_bytes(self, key: str, value: bytes): + path = os.path.join(self.base_dir, f"{key}") + os.makedirs(os.path.dirname(path), exist_ok=True) + with open(path, "wb") as f: + f.write(value) + + def get_bytes(self, key: str) -> bytes: + path = os.path.join(self.base_dir, f"{key}") + with open(path, "rb") as f: + return f.read() + + def get_value(self, key: str) -> str: + path = os.path.join(self.base_dir, f"{key}") + if not os.path.exists(path): + return "" + with open(os.path.join(self.base_dir, f"{key}"), "r") as f: + return f.read() + + def list_keys(self, key: str = "") -> List[str]: + path = os.path.join(self.base_dir, key) + if not os.path.exists(path): + return [] + # list all keys in the path + return [f for f in os.listdir(path)] + + def exists(self, key: str = "") -> bool: + path = os.path.join(self.base_dir, key) + return os.path.exists(path) + + +class SessionState(Enum): + IDLE = "idle" + WAITING = "waiting" + SKIPPED = "skipped" + SCHEDULED = "scheduled" + DONE = "done" + + +class ProfileRankState: + def __init__( + self, + store: LocalShmStore, + local_rank: int, + rank: int, + config: OnDemandProfilerConfig, + ) -> None: + self.store = store + self.local_rank = local_rank + self.rank = rank + self.config = config + self._update_state(SessionState.WAITING) + self.profiler: Optional[profiler.profile] = None + + # self._total_step = config.schedule_active + config.schedule_warmup + config.schedule_wait + if self.config.mode == "trace": + self._total_step = config.schedule_active + config.schedule_warmup + config.schedule_wait + else: + self._total_step = config.schedule_active + + def can_not_schedule(self, cur_step: int) -> bool: + return self.config.start_step > 0 and cur_step > self.config.start_step + + def step(self, cur_step: int) -> bool: + """ + Return True if profiling is done after step. + """ + if self.state == SessionState.WAITING: + if (self.config.start_step > 0 and cur_step == self.config.start_step) or ( + self.config.start_time is not None and datetime.now() >= self.config.start_time + ): + if self.config.profile_ranks[0] == -1 or self.rank in self.config.profile_ranks: + self._update_state(SessionState.SCHEDULED) + if self.config.mode == "trace": + # start dynamic profiler + self.start_profile() + logger.info( + f"Rank {self.rank}, local rank {self.local_rank}, " + f"Start Dynamic Profiler at {cur_step} step." + ) + else: + # start dump snapshot + self.start_dump() + logger.info( + f"Rank {self.rank}, local rank {self.local_rank}, " + f"Start Dump Snapshot at {cur_step} step." + ) + self.config.meta_data["start_step"] = cur_step + else: + self._update_state(SessionState.SKIPPED) + return False + + if self.state == SessionState.SCHEDULED or self.state == SessionState.SKIPPED: + self._total_step -= 1 + if self.config.mode == "trace": + if self.state == SessionState.SCHEDULED and self.profiler is not None: + self.profiler.step() + + if self._total_step == 0: + self._update_state(SessionState.DONE) + if self.profiler is not None: + self.profiler.stop() + self.profiler = None + return True + else: + return False + else: + # schedule dump + if self._total_step == 0: + self._update_state(SessionState.DONE) + # Only dump for profile ranks + if self.config.profile_ranks[0] == -1 or self.rank in self.config.profile_ranks: + self.stop_dump() + return True + else: + return False + + return True + + def _update_state(self, state: SessionState): + self.store.set_value(f"state_{self.local_rank}", state.value) + self.state = state + + def start_profile(self): + def trace_handler(prof): + try: + tb_handler = torch.profiler.tensorboard_trace_handler( + self.config.output_dir, + worker_name=( + f"trace_local_rank_{self.local_rank}_" + f"rank_{self.rank}_" + f"{socket.gethostname()}_{os.getpid()}" + ), + use_gzip=self.config.use_gzip, + ) + tb_handler(prof) + + if self.config.profile_memory: + memory_tl_file = os.path.join( + self.config.output_dir, + ( + f"memory_local_rank_{self.local_rank}_" + f"rank_{self.rank}_" + f"{socket.gethostname()}_{os.getpid()}.html" + ), + ) + logger.info("Export memory timeline to %s", memory_tl_file) + prof.export_memory_timeline(memory_tl_file) + except Exception as handler_error: + logger.error("Failed to handler trace profiler: %s", handler_error) + + try: + self.profiler = profiler.profile( + activities=self.config.activities, + schedule=profiler.schedule( + wait=self.config.schedule_wait, + warmup=self.config.schedule_warmup, + active=self.config.schedule_active, + ), + record_shapes=self.config.record_shapes, + profile_memory=self.config.profile_memory, + with_stack=self.config.with_stack, + with_flops=self.config.with_flops, + with_modules=self.config.with_modules, + on_trace_ready=trace_handler, + ) + except Exception as e: + logger.error("Failed to start profiler: %s", e) + self.profiler = None + + if self.profiler is not None: + self.profiler.start() + for key, value in self.config.meta_data.items(): + self.profiler.add_metadata_json(str(key), json.dumps(value)) + + def start_dump(self): + """ + Start memory dump snapshot + """ + + try: + logger.info("dump snapshot config %s", str(self.config)) + torch.cuda.memory._record_memory_history( + enabled=self.config.enabled, + context=self.config.context, + stacks=self.config.stacks, + max_entries=self.config.max_entries, + # device=self.config.device, + # clear_history=self.config.clear_history, + # compile_context=self.config.compile_context, + # global_record_annotations=self.config.global_record_annotations, + ) + except Exception as e: + logger.error("Failed to start dump snapshot: %s", e) + + def stop_dump(self): + """ + Save pickle and stop memory dump snapshot + """ + + try: + # save dump pickle + os.makedirs(self.config.output_dir, exist_ok=True) + dump_file = os.path.join( + self.config.output_dir, + ( + f"snap_local_rank_{self.local_rank}_" + f"rank_{self.rank}_" + f"{socket.gethostname()}_{os.getpid()}.pickle" + ), + ) + torch.cuda.memory._dump_snapshot(dump_file) + logger.info("Success to save pickle snapshot to %s", dump_file) + except Exception as save_error: + logger.error("Failed to save pickle snapshot: %s", save_error) + + # stop dump snapshot + try: + torch.cuda.memory._record_memory_history( + enabled=None, + ) + logger.info("Success to stop dump snapshot") + except Exception as enabled_error: + logger.error("Failed to stop dump snapshot: %s", enabled_error) + + +class ProfileSession: + def __init__( + self, + store: LocalShmStore, + config: OnDemandProfilerConfig, + local_world_size: int, + ) -> None: + self.store = store + self.session_id = config.session_id + self.config = config + self.local_world_size = local_world_size + + def get_rank_states(self) -> Dict[int, SessionState]: + states = {} + for lrank in range(self.local_world_size): + state = self.store.get_value(f"state_{lrank}") + if not state: + states[lrank] = SessionState.IDLE + else: + states[lrank] = SessionState(state) + return states + + def to_store(self): + self.store.set_bytes("config", pickle.dumps(self.config)) + self.store.set_value("local_world_size", str(self.local_world_size)) + + @classmethod + def from_store(cls, store: LocalShmStore) -> "ProfileSession": + config = pickle.loads(store.get_bytes("config")) + local_world_size = int(store.get_value("local_world_size")) + return cls(store, config, local_world_size) + + +class ProfilerNodeStateManager: + LOCAL_CONFIG_PATH = "/dev/shm/atorch_dp/local_config.json" + + def __init__(self) -> None: + self.store = LocalShmStore("atorch_dp") + + def get_dynamic_profile_enabled(self) -> bool: + return self.store.get_value("dynamic_profile_enabled") == "1" + + def set_dynamic_profile_enabled(self): + self.store.set_value("dynamic_profile_enabled", "1") + + def get_megatron_parallel_config(self) -> Optional[MegatronParallelConfig]: + megatron_parallel_config = self.store.get_value("megatron_parallel_config") + try: + return MegatronParallelConfig.from_dict(json.loads(megatron_parallel_config)) + except Exception as e: + logger.error(f"Failed to parse megatron parallel config from local shm store: {e}") + return None + + def set_megatron_parallel_config(self, config: MegatronParallelConfig): + self.store.set_value( + "megatron_parallel_config", + json.dumps(config.to_dict(), indent=4), + ) + + def create_session( + self, + config: OnDemandProfilerConfig, + ) -> ProfileSession: + store = LocalShmStore(f"atorch_dp/sessions/{config.session_id}") + return ProfileSession(store, config, local_world_size()) + + def get_session(self, session_id: str) -> Optional[ProfileSession]: + store = LocalShmStore(f"atorch_dp/sessions/{session_id}") + if not store.exists(): + return None + return ProfileSession.from_store(store) + + def profile_rank_state(self, local_rank: int, rank: int, config: OnDemandProfilerConfig) -> ProfileRankState: + store = LocalShmStore(f"atorch_dp/sessions/{config.session_id}") + return ProfileRankState(store, local_rank, rank, config) + + def list_sessions(self) -> List[ProfileSession]: + return [ + ProfileSession.from_store(LocalShmStore(f"atorch_dp/sessions/{session_id}")) + for session_id in self.store.list_keys("atorch_dp/sessions/") + ] + + def set_steps(self, steps: int): + self.store.set_value("steps", str(steps)) + + def get_steps(self) -> int: + return int(self.store.get_value("steps")) + + class _DynamicProfile(metaclass=SingletonMeta): def __init__(self) -> None: - self.profiler = None - self.config: Optional[ProfilerConfig] = None - self.step_num = 0 + self.cur_step = 0 self._optimizer_id = 0 - self._dynamic_monitor: Optional[ThreadFileConfigMonitor[ProfilerConfig]] = None + self._dynamic_monitor: Optional[ThreadFileConfigMonitor[OnDemandProfilerConfig]] = None + self._node_state_manager = ProfilerNodeStateManager() + + self._state: Optional[ProfileRankState] = None # hook step function if torch.__version__ >= "2.0.0": torch.optim.Optimizer._patch_step_function = _DynamicProfile.patch_step_function # type: ignore[assignment] @@ -115,80 +540,75 @@ def __init__(self) -> None: def init(self, cfg_path: str): logger.info("init dynamic profile with cfg path: %s", cfg_path) self.rank = rank() - self._dynamic_monitor = ThreadFileConfigMonitor(config_path=cfg_path, config_class=ProfilerConfig) + self.local_rank = local_rank() + self._dynamic_monitor = ThreadFileConfigMonitor( + config_paths=[ + self._node_state_manager.LOCAL_CONFIG_PATH, + cfg_path, + ], + config_class=OnDemandProfilerConfig, + ) self._dynamic_monitor.start() atexit.register(self._clean_resource) + if self.local_rank == 0: + # save megatron parallel config if available + megatron_config = MegatronParallelConfig.from_megatron_args() + if megatron_config is not None: + self._node_state_manager.set_megatron_parallel_config(megatron_config) + + # set dynamic profile enabled + self._node_state_manager.set_dynamic_profile_enabled() + + # start a thread to update steps every 1 minute + self._steps_updater = threading.Thread(target=self._update_steps, daemon=True) + self._steps_updater.start() + + def _update_steps(self): + while True: + self._node_state_manager.set_steps(self.cur_step) + time.sleep(60) def _clean_resource(self): - if self.profiler is not None: - self.profiler.stop() - self.profiler = None - logger.info("Profiler stop when process exit, check cfg json active whether over all step!") self._dynamic_monitor.stop() def step(self): self.cur_step += 1 - if self._dynamic_monitor is not None: - config = self._dynamic_monitor.get_config() - if config is not None: - self.config = config - - if self.profiler: - # step profiler - self.profiler.step() - self.step_num -= 1 - - # stop profiler if step num is 0 - if 0 == self.step_num: - self.profiler.stop() - self.profiler = None - - logger.info("Stop Dynamic Profiler at {} step.".format(self.cur_step)) - elif self.profiler is None and self.config is not None and self.cur_step == self.config.start_step: - # start profiler - self.step_num = self.config.schedule_active + self.config.schedule_warmup - self.start_profile() - self.config = None - - def start_profile(self): - def trace_handler(): - if self.config.profile_ranks[0] == -1 or self.rank in self.config.profile_ranks: - return torch.profiler.tensorboard_trace_handler( - self.config.output_dir, - worker_name=f"torch_profiler_rank_{self.rank}_{socket.gethostname()}_{os.getpid()}", - use_gzip=self.config.use_gzip, + if self._state is not None: + if self._state.step(self.cur_step): + self._state = None + logger.info( + f"Rank {self.rank}, local rank {self.local_rank}, " + f"Stop Dynamic Profiler or dump snapshot at {self.cur_step} step." ) - else: - logger.info("Profile will not be recorded for rank %d", self.rank) - - def _dummy_writer(p): - # Do nothing - pass - - return _dummy_writer - - self.profiler = profiler.profile( - activities=self.config.activities, - schedule=profiler.schedule( - wait=self.config.schedule_wait, - warmup=self.config.schedule_warmup, - active=self.config.schedule_active, - repeat=self.config.schedule_repeat, - skip_first=self.config.schedule_skip_first, - ), - record_shapes=self.config.record_shapes, - profile_memory=self.config.profile_memory, - with_stack=self.config.with_stack, - with_flops=self.config.with_flops, - with_modules=self.config.with_modules, - on_trace_ready=trace_handler(), - ) - self.profiler.start() - for key, value in self.config.meta_data.items(): - self.profiler.add_metadata_json(str(key), json.dumps(value)) - - logger.info("Start Dynamic Profiler at {} step.".format(self.cur_step)) + return + + if self._dynamic_monitor is None: + return + + # check if config is modified + config = self._dynamic_monitor.get_config_if_modified() + if config is None: + return + + # create profile rank state + state = self._node_state_manager.profile_rank_state(self.local_rank, self.rank, config) + if state.can_not_schedule(self.cur_step): + logger.info( + "start step is set to %d, but current step is %d, skip profiling", + config.start_step, + self.cur_step, + ) + return + + self._state = state + # store session info if local rank is 0 + if self.local_rank == 0: + session = self._node_state_manager.create_session(config) + session.to_store() + + # start profiler if start step or start time is set + self._state.step(self.cur_step) @staticmethod def patch_step_function(optimizer: torch.optim.Optimizer): diff --git a/atorch/utils/dynamic_profiler/_file_monitor.py b/atorch/utils/dynamic_profiler/_file_monitor.py index 29d74ce..9c5ad77 100644 --- a/atorch/utils/dynamic_profiler/_file_monitor.py +++ b/atorch/utils/dynamic_profiler/_file_monitor.py @@ -3,12 +3,14 @@ import threading import time from copy import deepcopy -from dataclasses import is_dataclass -from typing import Any, Callable, Dict, Generic, Optional, Type, TypeVar, cast +from dataclasses import field, is_dataclass +from datetime import datetime +from typing import Any, Callable, Dict, Generic, List, Optional, Tuple, Type, TypeVar, cast from atorch.common.log_utils import default_logger as logger T = TypeVar("T", bound=object) +_UPDATE_TIME_KEY = "x_config_update_time" def is_frozen_dataclass(config_class: Type[T]) -> bool: @@ -26,6 +28,11 @@ def is_frozen_dataclass(config_class: Type[T]) -> bool: return True +def datetime_field(date_format: str, default: Optional[datetime] = None, **kwargs) -> Any: + """Create a dataclass field for datetime objects with a specific format.""" + return field(default=default, metadata={"is_datetime": True, "format": date_format}, **kwargs) + + def create_dataclass_from_dict(config_dict: Dict[str, Any], config_class: Type[T]) -> Optional[T]: """Convert a dictionary to the specified config class instance.""" try: @@ -40,11 +47,21 @@ def create_dataclass_from_dict(config_dict: Dict[str, Any], config_class: Type[T elif isinstance(v, list): filtered_dict[k] = tuple(deepcopy(v)) - # support sub-dataclass for k, v in filtered_dict.items(): - k_type = config_class.__dataclass_fields__[k].type # type: ignore - if is_frozen_dataclass(k_type): - filtered_dict[k] = create_dataclass_from_dict(v, k_type) + k_field = config_class.__dataclass_fields__[k] # type: ignore + + # support sub-dataclass + if is_frozen_dataclass(k_field.type): + filtered_dict[k] = create_dataclass_from_dict(v, k_field.type) + + # support datetime + if k_field.metadata.get("is_datetime", False): + if isinstance(v, str): + filtered_dict[k] = datetime.strptime(v, k_field.metadata["format"]) + elif isinstance(v, datetime): + filtered_dict[k] = v + else: + raise ValueError(f"Invalid datetime value: {v}") # create dataclass instance and convert to T type instance = config_class(**filtered_dict) # type: ignore @@ -54,35 +71,71 @@ def create_dataclass_from_dict(config_dict: Dict[str, Any], config_class: Type[T return None +class FileWatcher: + def __init__(self, file_path: str, expire_time: int = 600): + self._file_path = file_path + self._expire_time = expire_time + self._last_mtime = 0.0 + + def has_changed(self) -> bool: + """Check if the file has changed and not expired.""" + if not os.path.exists(self._file_path): + return False + + current_mtime = os.path.getmtime(self._file_path) + if current_mtime > self._last_mtime: + self._last_mtime = current_mtime + # if expire time is set, check if the file is expired + if self._expire_time > 0: + if current_mtime + self._expire_time < time.time(): + return False + return True + return False + + def read_if_modified(self) -> Tuple[Optional[str], Optional[datetime]]: + """Read the file if it has been modified.""" + if self.has_changed(): + logger.info(f"File {self._file_path} has been modified within {self._expire_time} seconds") + with open(self._file_path, "r") as f: + return f.read(), datetime.fromtimestamp(self._last_mtime) + return None, None + + def __str__(self) -> str: + return f"FileWatcher(file_path={self._file_path}, expire_time={self._expire_time})" + + class ThreadFileConfigMonitor(Generic[T]): """ Generic ThreadFileConfigMonitor is used to monitor file changes and load the config into a specified immutable dataclass type. + + if all config files are modified, the last one will be used. """ def __init__( self, - config_path: str, + config_paths: List[str], config_class: Type[T], poll_interval: int = 60, + expire_time: int = 600, validator: Optional[Callable[[T], bool]] = None, ): """ Initialize the file monitor. Args: - config_path: The path of the configuration file. + config_paths: The path of the configuration files. config_class: The dataclass type to load the config into. poll_interval: The polling interval (seconds). validator: Optional function to validate the loaded config. """ - self._config_path = config_path - + self._file_monitors = [FileWatcher(path, expire_time) for path in config_paths] if not is_frozen_dataclass(config_class): raise TypeError(f"{config_class} must be a frozen dataclass") self._config_class: Type[T] = config_class # type: ignore self._poll_interval = poll_interval + self._expire_time = expire_time self._validator = validator self._last_mtime: float = 0.0 self._current_config: Optional[T] = None @@ -90,6 +143,8 @@ def __init__( self._thread = None self._lock = threading.Lock() + self._check_config_if_modified = False + def start(self): """Start the monitor thread.""" if self._thread is not None and self._thread.is_alive(): @@ -107,22 +162,24 @@ def stop(self): def _monitor_loop(self): """Monitor loop, check file changes.""" - logger.info(f"Start monitoring file {self._config_path}") + logger.info("Start monitoring files: %s", self._file_monitors) while self._running: try: - if self._check_file_changed(): - logger.info(f"File {self._config_path} has changed") - config_dict = self._read_config() - if config_dict: - # convert dict to immutable dataclass - config = create_dataclass_from_dict(config_dict, self._config_class) - # validate config - if config is not None and self._is_valid_config(config): - with self._lock: - logger.info(f"Update config {config}") - self._current_config = config - else: - logger.warning(f"Invalid configuration found in {self._config_path}") + for file_monitor in self._file_monitors: + content, last_mtime = file_monitor.read_if_modified() + if content: + config_dict = json.loads(content) + if config_dict: + # inject update time + config_dict[_UPDATE_TIME_KEY] = last_mtime + # convert dict to immutable dataclass + config = create_dataclass_from_dict(config_dict, self._config_class) + # validate config + if config is not None and self._is_valid_config(config): + with self._lock: + logger.info(f"Update config {config}") + self._current_config = config + self._check_config_if_modified = True except Exception as e: logger.error(f"Error in file monitor: {e}") @@ -138,26 +195,6 @@ def _is_valid_config(self, config: T) -> bool: return True - def _check_file_changed(self) -> bool: - """Check if the file has changed.""" - if not os.path.exists(self._config_path): - return False - - current_mtime = os.path.getmtime(self._config_path) - if current_mtime > self._last_mtime: - self._last_mtime = current_mtime - return True - return False - - def _read_config(self) -> Dict[str, Any]: - """Read the configuration file content.""" - try: - with open(self._config_path, "r") as f: - return json.load(f) - except Exception as e: - logger.error(f"Error reading config file {self._config_path}: {e}") - return {} - def get_config(self) -> Optional[T]: """ Get the current configuration. @@ -169,6 +206,16 @@ def get_config(self) -> Optional[T]: with self._lock: return self._current_config + def get_config_if_modified(self) -> Optional[T]: + """ + Get the current configuration if it has been modified after last call. + """ + with self._lock: + if self._current_config is not None and self._check_config_if_modified: + self._check_config_if_modified = False + return self._current_config + return None + def set_poll_interval(self, seconds: int): """Set the polling interval.""" if seconds > 0: diff --git a/atorch/utils/inspector/hooks.py b/atorch/utils/inspector/hooks.py index 6c7e11b..ee98c17 100644 --- a/atorch/utils/inspector/hooks.py +++ b/atorch/utils/inspector/hooks.py @@ -1,3 +1,4 @@ +import inspect import os import re @@ -85,6 +86,7 @@ def __init__( except OSError: logger.error(f"Cannot create directory {plot_tensor_dir}, plot_tensor is disabled.") self.plot_tensor = False + self.hooks = [] def enable(self, if_enable): self.enable = if_enable @@ -102,6 +104,7 @@ def register_hooks( exclude_tensor_name_pattern=None, layer_types=(torch.nn.Linear,), backward_use_e4m3=False, + te_fp8_check=False, ): """Register log tensor hook and save tensor hook""" @@ -120,25 +123,38 @@ def register_hooks( ): # log tensor hook - layer.register_forward_hook(log_tensor_hook(name, self, is_fwd=True)) - layer.register_full_backward_hook( - log_tensor_hook(name, self, is_fwd=False, backward_use_e4m3=backward_use_e4m3) + hook = layer.register_forward_hook(log_tensor_hook(name, self, is_fwd=True, te_fp8_check=te_fp8_check)) + self.hooks.append(hook) + hook = layer.register_full_backward_hook( + log_tensor_hook( + name, self, is_fwd=False, backward_use_e4m3=backward_use_e4m3, te_fp8_check=te_fp8_check + ) ) + self.hooks.append(hook) # save tensor hook if self.save_tensor: - layer.register_forward_hook(save_tensor_hook(name, self, is_fwd=True)) - layer.register_full_backward_hook(save_tensor_hook(name, self, is_fwd=False)) + hook = layer.register_forward_hook(save_tensor_hook(name, self, is_fwd=True)) + self.hooks.append(hook) + hook = layer.register_full_backward_hook(save_tensor_hook(name, self, is_fwd=False)) + self.hooks.append(hook) # plot tensor hook if self.plot_tensor: - layer.register_forward_hook(plot_tensor_hook(name, self, is_fwd=True)) - layer.register_full_backward_hook(plot_tensor_hook(name, self, is_fwd=False)) + hook = layer.register_forward_hook(plot_tensor_hook(name, self, is_fwd=True)) + self.hooks.append(hook) + hook = layer.register_full_backward_hook(plot_tensor_hook(name, self, is_fwd=False)) + self.hooks.append(hook) matched_modules.append(name) return matched_modules + def remove_hooks(self): + for hook in self.hooks: + hook.remove() + self.hooks = [] + def calculate_percentiles(activations): min_vals = np.min(activations, axis=0) @@ -260,7 +276,10 @@ def hook(module, inputs, outputs): else: tensors = {} - tensors[tensor_names[1]] = outputs[0].detach().cpu() # y or dy + if isinstance(outputs, tuple): + tensors[tensor_names[1]] = outputs[0].detach().cpu() # y or dy + else: + tensors[tensor_names[1]] = outputs.detach().cpu() # save weight if is_fwd: @@ -290,7 +309,7 @@ def hook(module, inputs, outputs): return hook -def log_tensor_hook(module_name, inspector, is_fwd=True, backward_use_e4m3=False): +def log_tensor_hook(module_name, inspector, is_fwd=True, backward_use_e4m3=False, te_fp8_check=False): """Set up hook for forward or backpropagation""" interval = inspector.log_tensor_interval log_fn = inspector.log_fn @@ -322,7 +341,7 @@ def hook(module, inputs, outputs): # get target tensor: (fwd_x & fwd_w) or bwd_dy targets = [ - inputs[0] if is_fwd else outputs[0], + inputs[0] if is_fwd else outputs[0] if isinstance(outputs, tuple) else outputs, ] if is_fwd: for p_name, p_tensor in module.named_parameters(recurse=False): @@ -420,6 +439,12 @@ def hook(module, inputs, outputs): log_str += f", blockwise_underflows(%): {results['blockwise']:.1f}" tb_dict["fp8_blockwise_underflows"] = results["blockwise"] + if te_fp8_check and index == 0: + fp8_compute_cos = get_te_fp8_compute_cos_similarity(module, inputs, outputs, is_fwd) + if fp8_compute_cos is not None: + log_str += f", fp8 compute cos similarity: {fp8_compute_cos:.3e}" + tb_dict["fp8_compute_cos_similarity"] = fp8_compute_cos + log_fn(log_str) if inspector.summary_writer is not None: # write to tensorboard @@ -457,7 +482,7 @@ def hook(module, inputs, outputs): if inspector.enable and (step % interval) == 0 and step != 0: # get target tensor: (fwd_x & fwd_w) or bwd_dy targets = [ - inputs[0] if is_fwd else outputs[0], + inputs[0] if is_fwd else outputs[0] if isinstance(outputs, tuple) else outputs, ] # TODO: enable weight later on @@ -480,3 +505,84 @@ def hook(module, inputs, outputs): log_fn(f"[plot tensor hook] Save tensor plot figure: {filename}") return hook + + +def get_te_fp8_compute_cos_similarity(module, inputs, outputs, is_fwd): + try: + from transformer_engine.pytorch.cpp_extensions import general_grouped_gemm + from transformer_engine.pytorch.module.base import _2X_ACC_DGRAD, get_multi_stream_cublas_workspace + from transformer_engine.pytorch.module.grouped_linear import GroupedLinear, _GroupedLinear + from transformer_engine.pytorch.tensor.quantized_tensor import QuantizedTensorBase + except (ImportError, ModuleNotFoundError): + return None + + if isinstance(module, GroupedLinear): + weight_tensors = [getattr(module, f"weight{i}") for i in range(module.num_gemms)] + if not isinstance(weight_tensors[0], QuantizedTensorBase): + return None + bias_tensors = [getattr(module, f"bias{i}") for i in range(module.num_gemms)] + weight_tensors = [w.dequantize() for w in weight_tensors] + # Assuming only weights in fp8 format, inputs/outputs are in non-fp8 formats. + if is_fwd: + # bf16 forward to get bf16_outputs + none_quantizers = [None] * module.num_gemms + args = [ + None, + inputs[0], + inputs[1], + module.apply_bias, + False, # is_first_microbatch + False, # fp8 + False, # fp8_calibration + None, # wgrad_store + none_quantizers, # input_quantizers + none_quantizers, # weight_quantizers + none_quantizers, # output_quantizers + none_quantizers, # grad_output_quantizers + False, # fuse_wgrad_accumulation + False, # is_cpu_offload_enabled + module.sequence_parallel, + inputs[0].dtype, # activation_dtype + False, # is_grad_enabled + module, + None, # skip_fp8_weight_update + ] + # new te version adds save_original_input param. + sig = inspect.signature(_GroupedLinear.forward) + if "save_original_input" in sig.parameters.keys(): + args.append(False) + args += [ + *weight_tensors, + *bias_tensors, + ] + bf16_outputs = _GroupedLinear.forward(*args) + module._cache_m_splits = inputs[1] + cos_similarity = cosine(bf16_outputs, outputs[0] if isinstance(outputs, tuple) else outputs) + else: + # bf16 backward to get bf16_outputs + g_output = outputs[0] if isinstance(outputs, tuple) else outputs + grad_output = g_output.contiguous() + grad_output_view = grad_output.view(-1, grad_output.shape[-1]) + grad_output_mats = torch.split(grad_output_view, module._cache_m_splits) + bf16_outputs = torch.empty( + sum(module._cache_m_splits), + weight_tensors[0].shape[1], + dtype=g_output.dtype, + device=g_output.device, + ) + general_grouped_gemm( + weight_tensors, + grad_output_mats, + [bf16_outputs], + g_output.dtype, + get_multi_stream_cublas_workspace(), + single_output=True, + layout="NN", + m_splits=module._cache_m_splits, + grad=True, + use_split_accumulator=_2X_ACC_DGRAD, + ) + module._cache_m_splits = None + cos_similarity = cosine(bf16_outputs, inputs[0]) + return cos_similarity + return None diff --git a/atorch/utils/parse_memory_pickle.py b/atorch/utils/parse_memory_pickle.py new file mode 100644 index 0000000..0f2a8e7 --- /dev/null +++ b/atorch/utils/parse_memory_pickle.py @@ -0,0 +1,125 @@ +from __future__ import absolute_import, unicode_literals + +import glob +import pickle +from argparse import ArgumentParser +from collections import Counter + +from tqdm import tqdm + + +def filename_is_allow(filename, whitelist): + if not filename.endswith(".py"): + return False + # return True + for whiteitem in whitelist: + if whiteitem in filename: + return True + return False + + +def parse_memory_pickle_file(pickle_file, blocklist, whitelist): + with open(pickle_file, "rb") as f: + pickle_obj = pickle.load(f) + filename_counter = Counter() + max_memory_filename_counter = Counter() + addr2filename = {} + total_memory = 0 + max_memory = 0 + for device_trace in pickle_obj["device_traces"]: + + for malloc_record in tqdm(device_trace): + # # {'alloc', 'free_completed', 'free_requested', 'segment_map', 'segment_unmap'} + if malloc_record["action"] == "alloc": + total_memory += malloc_record["size"] + frames = malloc_record["frames"] + frame_filenames = [] + # only record the first file in whitelist + has_record = False + for frame in frames: + if filename_is_allow(frame["filename"], whitelist): + filename_lineno = frame["filename"] + ":" + str(frame["line"]) + if not has_record: + filename_counter[filename_lineno] += malloc_record["size"] + addr2filename[malloc_record["addr"]] = filename_lineno + has_record = True + + frame_filenames.append(filename_lineno) + + elif malloc_record["action"] == "free_completed": + total_memory -= malloc_record["size"] + frames = malloc_record["frames"] + frame_filenames = [] + has_record = False + if malloc_record["addr"] in addr2filename: + filename_lineno = addr2filename.pop(malloc_record["addr"]) + filename_counter[filename_lineno] -= malloc_record["size"] + has_record = True + elif malloc_record["action"] == "free_requested": + pass + # total_memory -= malloc_record["size"] + elif malloc_record["action"] == "segment_map": + pass + # total_memory += malloc_record["size"] + elif malloc_record["action"] == "segment_unmap": + pass + # total_memory -= malloc_record["size"] + if total_memory > max_memory: + max_memory = total_memory + for name, memory_alloc in filename_counter.items(): + max_memory_filename_counter[name] = memory_alloc + print(f"{pickle_file} max_memory: {max_memory/1024/1024:.2f}MB") + print(f"final total_memory {total_memory/1024/1024:.2f}MB") + + return max_memory_filename_counter, filename_counter + + +def print_result(filename_counter, blocklist, limit=20): + MB = 1024 * 1024 + print("=" * 100) + show_count = 0 + for filename, memory_alloc in filename_counter.most_common(): + skip = False + for blockname in blocklist: + if filename.startswith(blockname): + skip = True + break + if skip: + continue + show_count += 1 + print(filename, f"alloc{memory_alloc/MB:.2f}MB") + if show_count > limit: + break + + +def parse_args(argv=None): + parser = ArgumentParser(usage="""python parse_memory_pickle.py memory_pickle_file""") + + parser.add_argument("pickle_files", nargs="*") + + parser.add_argument("--verbose", "-v", action="store_true") + parser.add_argument("--blocklist", nargs="*") + parser.add_argument("--whitelist", nargs="*", default=["torch", "atorch", "megatron", "transformer_engine"]) + parser.add_argument("--show_count", type=int, default=20) + + args = parser.parse_args(argv) + return args + + +def main(): + args = parse_args() + # TODO: multiprocess + if len(args.pickle_files) == 1 and "*" in args.pickle_files[0]: + filenames = glob.glob(args.pickle_files[0]) + else: + filenames = args.pickle_files + + blocklist = args.blocklist or [] + whitelist = args.whitelist or [] + for filename in filenames: + filename_counter, final_counter = parse_memory_pickle_file(filename, whitelist) + print_result(filename_counter, blocklist, whitelist) + + +if __name__ == "__main__": + main() diff --git a/atorch/utils/parse_trace_json.py b/atorch/utils/parse_trace_json.py index 67a481c..ff1003f 100644 --- a/atorch/utils/parse_trace_json.py +++ b/atorch/utils/parse_trace_json.py @@ -7,15 +7,18 @@ """ from __future__ import absolute_import, unicode_literals +import glob from argparse import ArgumentParser +from datetime import datetime from json import load +from multiprocessing import Pool import numpy as np import pandas as pd def get_compute_kernel(df): - return df.query("cat == 'kernel' and ~name.str.startswith('nccl')") + return df.query("cat == 'kernel' and ~name.str.startswith('pccl') and ~name.str.startswith('nccl')") def prepare_df(json_obj): @@ -61,7 +64,9 @@ def analyze_gpu_kernel(df, verbose=False): "div": "DivFunctor", } op_times = { - "gemm": gpu_kernel_df[gpu_kernel_df.name.str.contains("gemm")]["dur"].sum(), + "gemm": gpu_kernel_df[gpu_kernel_df.name.str.contains("gemm") ^ gpu_kernel_df.name.str.contains("nvjet")][ + "dur" + ].sum(), "elementwise": { "total": gpu_kernel_df[gpu_kernel_df.name.str.contains("elementwise")]["dur"].sum(), }, @@ -130,9 +135,13 @@ def fused_kernels(kernels): def analyze_communicate_overlap(df, verbose=False): - comm_kernel_df = df.query("name.str.startswith('ncclKernel')|name.str.startswith('ncclDevKernel')") + comm_kernel_df = df.query( + "name.str.startswith('ncclKernel')|name.str.startswith('ncclDevKernel')|name.str.startswith('pcclKernel')" + ) all_comm_tids = list(set(comm_kernel_df["tid"].values)) - all_compute_tids = list(set(df.query("~name.str.startswith('nccl')")["tid"].values)) + all_compute_tids = list( + set(df.query("~name.str.startswith('nccl') and ~name.str.startswith('pccl')")["tid"].values) + ) if verbose: print("all_comm_tids", all_comm_tids) print("all_compute_tids", all_compute_tids) @@ -146,6 +155,10 @@ def analyze_communicate_overlap(df, verbose=False): "nooverlap_comm_df": comm_kernel_df, } comm_time_us = comm_kernel_df["dur"].sum() + all2all_time_us = comm_kernel_df[comm_kernel_df["name"].str.contains("SendRecv")]["dur"].sum() + fsdp_time_us = comm_kernel_df.query("name.str.contains('AllGather') or name.str.contains('ReduceScatter')")[ + "dur" + ].sum() # int64 = int64 + int64, can dur convert to int64? # gpu_kernel_df.loc[:,"finish_time"] = gpu_kernel_df["ts"] + gpu_kernel_df["dur"] # comm_kernel_df.loc[:,"finish_time"] = comm_kernel_df["ts"] + comm_kernel_df["dur"] @@ -238,47 +251,106 @@ def analyze_communicate_overlap(df, verbose=False): "comm_time_us": comm_time_us, "overlap_time_us": overlap_time, "nooverlap_comm_df": nooverlap_comm_df, + "all2all_time_us": all2all_time_us, + "fsdp_time_us": fsdp_time_us, } +def parse_trace_file(json_file, verbose=False): + with open(json_file, "r") as fin: + json_obj = load(fin) # TODO: iter json_obj,save memory usage + df = prepare_df(json_obj) + min_ts = df["ts"].min() + if "profiler_starttime" in json_obj: + # convert to abs time + df["ts"] = df["ts"] + json_obj["profiler_starttime"] - min_ts + df["finish_time"] = df["finish_time"] + json_obj["profiler_starttime"] - min_ts + kernel_start_time = df.query("cat=='kernel'")["ts"].min() + # print("kernel 5 sample:\n", df.query("cat=='kernel'").head(n=5)) + + summary = analyze_gpu_kernel(df) + ret_communicate = analyze_communicate_overlap(df, verbose) + ret_communicate.pop("nooverlap_comm_df") + # print("communicate summary:", ret_communicate) + summary.update(ret_communicate) + # distributedInfo: {'backend': 'nccl', 'rank': 10, 'world_size': 16} + if "distributedInfo" in json_obj: + rank = json_obj["distributedInfo"]["rank"] + else: + rank = 0 + summary["rank"] = rank + prepare_sample_df = df.query( + "cat=='python_function' and name.str.contains('torch/utils/data/dataloader.py')" + " and name.str.contains('__next__')" + ) + _train_inner_loop_time = df.query("cat=='python_function' and name.str.contains('_train_inner_loop')") + backward_time = df.query("cat=='python_function' and name.str.contains('_engine_run_backward')")["dur"].sum() + summary["sample_time"] = int(prepare_sample_df["dur"].sum()) + summary["total_time"] = int(_train_inner_loop_time["dur"].sum()) + summary["backward_time"] = int(backward_time) + summary["gpu_time_us"] = int(summary["gpu_time_us"]) + summary["comm_time_us"] = int(summary["comm_time_us"]) + summary["overlap_time_us"] = int(summary["overlap_time_us"]) + summary["memory_time"] = summary["op_cat_times"]["memory"] + return summary, kernel_start_time + + +def print_profiler_summary(all_summary, kernel_start_times, all_summary_path=None): + all_summary_df = pd.DataFrame(all_summary) + print(all_summary_df) + if all_summary_path: + with open(all_summary_path, "w") as fout: + all_summary_df.to_csv(fout, index=False) + # check if all kernels start at the same time + pd.set_option("display.float_format", "{:.3f}".format) + print(all_summary_df.iloc[all_summary_df["comm_time_us"].idxmin()].op_cat_times) + print( + "total_time min/max rank is", + all_summary_df.iloc[all_summary_df["comm_time_us"].idxmin()], + all_summary_df.iloc[all_summary_df["comm_time_us"].idxmax()], + ) + print(all_summary_df.describe()) + kernel_start_times = np.asarray(kernel_start_times) + + min_start_time = np.min(kernel_start_times) + print( + "min start_time=", + datetime.fromtimestamp(min_start_time / 1e6), + "max start_time=", + datetime.fromtimestamp(np.max(kernel_start_times) / 1e6), + ) + print(kernel_start_times - min_start_time) + + def main(): parser = ArgumentParser(usage="""python parse_trace_json.py trace_1.json""") parser.add_argument( "--all_summary_path", ) + parser.add_argument("--process_num", "-p", type=int, default=1) parser.add_argument("json_files", nargs="*") + parser.add_argument("--verbose", "-v", action="store_true") args = parser.parse_args() kernel_start_times = [] all_summary = [] - for json_file in args.json_files: - with open(json_file, "r") as fin: - json_obj = load(fin) # TODO: iter json_obj,save memory usage - df = prepare_df(json_obj) - kernel_start_times.append(df.query("cat=='kernel'")["ts"].min()) - print("kernel 5 sample:\n", df.query("cat=='kernel'").head(n=5)) - - ret = analyze_gpu_kernel(df) - print("compute summary:", ret) - ret_communicate = analyze_communicate_overlap(df, args.verbose) - nooverlap_comm_df = ret_communicate.pop("nooverlap_comm_df") - print("communicate summary:", ret_communicate) - ret.update(ret_communicate) - # distributedInfo: {'backend': 'nccl', 'rank': 10, 'world_size': 16} - rank = json_obj["distributedInfo"]["rank"] - ret["rank"] = rank - all_summary.append(ret) - print("no overlap comm op:\n", nooverlap_comm_df) - all_summary_df = pd.DataFrame(all_summary) - print(all_summary_df) - if args.all_summary_path: - with open(args.all_summary_path, "w") as fout: - all_summary_df.to_csv(fout, index=False) - # check if all kernels start at the same time - kernel_start_times = np.asarray(kernel_start_times) - min_start_time = np.min(kernel_start_times) - print(kernel_start_times - min_start_time) + if len(args.json_files) == 1 and "*" in args.json_files[0]: + filenames = glob.glob(args.json_files[0]) + else: + filenames = args.json_files + if args.process_num > 1: + pool = Pool(args.process_num) + multi_result = pool.map(parse_trace_file, filenames) + for summary, kernel_start_time in multi_result: + all_summary.append(summary) + kernel_start_times.append(kernel_start_time) + else: + for json_file in filenames: + summary, kernel_start_time = parse_trace_file(json_file) + all_summary.append(summary) + kernel_start_times.append(kernel_start_time) + print_profiler_summary(all_summary, kernel_start_times) if __name__ == "__main__": diff --git a/atorch/utils/virtual_optimizer/megatron_virtual_optimizer.py b/atorch/utils/virtual_optimizer/megatron_virtual_optimizer.py index 99cdc9a..4def450 100644 --- a/atorch/utils/virtual_optimizer/megatron_virtual_optimizer.py +++ b/atorch/utils/virtual_optimizer/megatron_virtual_optimizer.py @@ -3,18 +3,20 @@ import torch from atorch.common.log_utils import default_logger as logger -from atorch.trainer.args import AtorchTrainingArgs + +# from atorch.trainer.args import AtorchTrainingArgs # noqa: E402 from atorch.utils.import_util import is_megatron_lm_available +from atorch.utils.virtual_optimizer.patch_utils import ( + patch_chained_optimizer, + patch_distributed_optimizer, + virtual_distributed_optimizer_load_state_dict, + zero_out_shard_fp32_memory, +) if is_megatron_lm_available(): from megatron.core.dist_checkpointing.mapping import ShardedStateDict - from megatron.core.optimizer import ( - ChainedOptimizer, - MegatronOptimizer, - _get_param_groups, - _update_min_and_max_lr_in_param_groups, - ) - from megatron.core.optimizer.distrib_optimizer import MixedPrecisionOptimizer + from megatron.core.optimizer import ChainedOptimizer, MegatronOptimizer, _get_param_groups + from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer, MixedPrecisionOptimizer from megatron.core.optimizer.grad_scaler import ConstantGradScaler, DynamicGradScaler, MegatronGradScaler from megatron.core.optimizer.optimizer_config import OptimizerConfig @@ -124,9 +126,15 @@ def init_state_fn(opt): return optimizer +class VirtualChainedOptimizer(ChainedOptimizer): + @torch.no_grad() + def step(self): + return True, 0.0, 0 + + # Configure use_virtual_optimizer: true in antllm yaml file to enable this feature -def get_megatron_virtual_optimizer( - train_args: AtorchTrainingArgs, +def get_megatron_virtual_optimizer_v1( + train_args: Any, config: OptimizerConfig, model_chunks: List[MegatronModule], no_weight_decay_cond: Optional[Callable] = None, @@ -142,20 +150,14 @@ def get_megatron_virtual_optimizer( no_weight_decay_cond, scale_lr_cond, lr_mult, - use_decoupled_learning_rate=config.decoupled_lr is not None, - ) - - param_groups = _update_min_and_max_lr_in_param_groups( - param_groups, lr=config.lr, min_lr=config.min_lr, decoupled_lr=config.decoupled_lr, decoupled_min_lr=config.decoupled_min_lr, + # use_decoupled_learning_rate=config.decoupled_lr is not None, ) - dense_param_groups = list(filter(lambda g: not g["is_expert_parallel"], param_groups)) moe_param_groups = list(filter(lambda g: g["is_expert_parallel"], param_groups)) - optimizers = [ _get_megatron_virtual_optimizer_based_on_param_groups( config, @@ -173,20 +175,20 @@ def get_megatron_virtual_optimizer( if len(optimizers) == 1: return optimizers[0] - return ChainedOptimizer(optimizers) + opt = VirtualChainedOptimizer(optimizers) + return opt class MegatronVirtualOptimizer(MixedPrecisionOptimizer): def __init__( self, - # train_args: AtorchTrainingArgs, optimizer: torch.optim.Optimizer, config: OptimizerConfig, grad_scaler: MegatronGradScaler, init_state_fn: Optional[Callable] = None, ): super().__init__(optimizer, config, grad_scaler, init_state_fn) - # self.train_args = train_args + self.is_stub_optimizer = False if hasattr(self.optimizer, "param_groups_master"): self.param_groups_master: List[Any] = [] @@ -231,3 +233,43 @@ def zero_grad(self, set_to_none=True): # for groups in (): # for group in groups: # _zero_grad_group_helper(group, set_to_none) + + def _copy_model_grads_to_main_grads(self): + return + + def _copy_main_params_to_model_params(self): + return + + +def get_megatron_virtual_optimizer( + train_args: Any, + config: OptimizerConfig, + model_chunks: List[MegatronModule], + no_weight_decay_cond: Optional[Callable] = None, + scale_lr_cond: Optional[Callable] = None, + lr_mult: float = 1.0, +): + from megatron.core.optimizer import get_megatron_optimizer + + setattr( + DistributedOptimizer, + "load_state_dict", + staticmethod(virtual_distributed_optimizer_load_state_dict), + ) + + org_chained_optimizer = get_megatron_optimizer( + config, + model_chunks, + no_weight_decay_cond, + scale_lr_cond, + lr_mult, + ) + + zero_out_shard_fp32_memory(org_chained_optimizer) + + patch_chained_optimizer(org_chained_optimizer) + + for opt in org_chained_optimizer.chained_optimizers: + patch_distributed_optimizer(opt) + + return org_chained_optimizer diff --git a/atorch/utils/virtual_optimizer/patch_utils.py b/atorch/utils/virtual_optimizer/patch_utils.py new file mode 100644 index 0000000..e26b87a --- /dev/null +++ b/atorch/utils/virtual_optimizer/patch_utils.py @@ -0,0 +1,176 @@ +import types +from typing import List + +import torch + +from atorch.utils.import_util import is_megatron_lm_available + +if is_megatron_lm_available(): + print("********** megatron_lm_available: true **********") + from megatron.core.dist_checkpointing.mapping import ShardedStateDict + from megatron.core.tensor_parallel import param_is_not_tensor_parallel_duplicate + from megatron.core.transformer.module import param_is_not_shared +else: + print("********** megatron_lm_available: false **********") + from typing import Any, Dict + + ShardedStateDict = Dict[str, Any] + + def param_is_not_tensor_parallel_duplicate(param): + """Returns true if the passed-in parameter is not a duplicate parameter + on another TP rank.""" + return hasattr(param, "tensor_model_parallel") and param.tensor_model_parallel + + def param_is_not_shared(param): + return not hasattr(param, "shared") or not param.shared + + +def zero_out_shard_fp32_memory(chained_optimizer): + """ + Set the memory of all params in shard_fp32_groups and shard_fp32_from_float16_groups + of each DistributedOptimizer in chained_optimizers to zero, but keep the objects. + """ + print("********** zero_out_shard_fp32_memory **********") + for opt in chained_optimizer.chained_optimizers: + # DistributedOptimizer + if not hasattr(opt, "shard_fp32_groups") or not hasattr(opt, "shard_fp32_from_float16_groups"): + continue + # shard_fp32_groups + for group in opt.shard_fp32_groups: + for param in group: + if isinstance(param, torch.Tensor): + param.data = torch.empty((1,), dtype=param.dtype, device=param.device) + # shard_fp32_from_float16_groups + for group in opt.shard_fp32_from_float16_groups: + for param in group: + if isinstance(param, torch.Tensor): + param.data = torch.empty((1,), dtype=param.dtype, device=param.device) + + +def patch_distributed_optimizer(distributed_optimizer): + def virtual_get_main_grads_for_grad_norm(self) -> List[torch.Tensor]: + """ + Get main_grads that should be taken into account to compute the grad norm. + Filter parameters based on: + - grad should not be None. + - parameter should not be shared (i.e., grads shouldn't be double counted while + computing norms). + - should not be a replica due to tensor model parallelism. + """ + print("********** virtual_distributed_optimizer get_main_grads_for_grad_norm**********") + + params = self.get_parameters() + grads_for_norm = [] + for param in params: + if self.config.use_precision_aware_optimizer: + grad = param.decoupled_grad if hasattr(param, "decoupled_grad") else None + else: + grad = param.virtual_grad + grad_not_none = grad is not None + is_not_shared = param_is_not_shared(param) + is_not_tp_duplicate = param_is_not_tensor_parallel_duplicate(param) + + if grad_not_none and is_not_shared and is_not_tp_duplicate: + grads_for_norm.append(grad) + + return grads_for_norm + + def virtual_copy_model_grads_to_main_grads(self): + """ + Copy model grads to main grads. + + Since this step follows a reduce-scatter through the DDP's grad + buffer, this method is responsible for copying the updated grads + from the grad buffer to the main shard's grad field. + """ + print("********** virtual_distributed_optimizer copy model grads to main grads **********") + if self.is_stub_optimizer: + return + + # Utility method for copying group grads. + def copy_group_grads(model_groups, shard_main_groups): + for model_group, shard_main_group in zip(model_groups, shard_main_groups): + for model_param, shard_main_param in zip(model_group, shard_main_group): + + param_range_map = self._get_model_param_range_map(model_param) + param_range = param_range_map["param"] + # assert param_range.size == shard_main_param.nelement() + + model_grad = model_param.main_grad + shard_model_grad = model_grad.view(-1)[param_range.start : param_range.end] + if self.config.use_precision_aware_optimizer: + # Pytorch requires a param and its' grad to be the same dtype, but we want + # their types to be different in precision-aware optimizer. So we use + # ".decoupled_grad" to replace ".grad". + # Note that this requires corresponding modifications in the optimizer (Let + # the optimizer read gradients from ".decoupled_grad" instead of ".grad"). + shard_main_param.decoupled_grad = shard_model_grad + else: + # use grad will cause assgin error because grad is zeroed out in zero_out_shard_fp32_memory + shard_main_param.virtual_grad = shard_model_grad.float() + + # Copy model groups to shard groups. + if self.config.use_precision_aware_optimizer: + copy_group_grads(self.model_float16_groups, self.shard_float16_groups) + copy_group_grads(self.model_fp32_groups, self.shard_fp32_groups) + else: + copy_group_grads(self.model_float16_groups, self.shard_fp32_from_float16_groups) + copy_group_grads(self.model_fp32_groups, self.shard_fp32_groups) + + distributed_optimizer.get_main_grads_for_grad_norm = types.MethodType( + virtual_get_main_grads_for_grad_norm, distributed_optimizer + ) + + distributed_optimizer._copy_model_grads_to_main_grads = types.MethodType( + virtual_copy_model_grads_to_main_grads, distributed_optimizer + ) + + +def patch_chained_optimizer(chained_optimizer): + @torch.no_grad() + def virtual_step(self): + print("********** virtual_chained_optimizer step**********") + + return True, 0.0, 0 + + def virtual_load_state_dict(self, state_dict): + print("********** virtual_chained_optimizer load_state_dict**********") + + pass + + def virtual_sharded_state_dict( + self, + model_sharded_state_dict: ShardedStateDict, + is_loading: bool = False, + sharding_type: str = "fully_sharded_model_space", + ): + print("********** virtual_chained_optimizer sharded_state_dict**********") + + return {} + + def virtual_reload_model_params(self): + """Refreshes any internal state from the current model parameters. + Call whenever the parameters are changed outside of the optimizer. + For example, when we load a model from a checkpoint without loading + the optimizer, the model parameters are updated but for fp16 optimizer + with main parameters, the main parameters need to also be updated.""" + print("********** virtual_chained_optimizer reload_model_params**********") + + pass + + def virtual_load_parameter_state(self, filename: str, *, update_legacy_format: bool = False): + print("********** virtual_chained_optimizer load_parameter_state**********") + + pass + + chained_optimizer.step = types.MethodType(virtual_step, chained_optimizer) + chained_optimizer.load_state_dict = types.MethodType(virtual_load_state_dict, chained_optimizer) + chained_optimizer.sharded_state_dict = types.MethodType(virtual_sharded_state_dict, chained_optimizer) + chained_optimizer.reload_model_params = types.MethodType(virtual_reload_model_params, chained_optimizer) + chained_optimizer.load_parameter_state = types.MethodType(virtual_load_parameter_state, chained_optimizer) + + +def virtual_distributed_optimizer_load_state_dict(self, state_dict): + print("********** virtual_distributed_optimizer load_state_dict **********") + + pass diff --git a/atorch/utils/virtual_optimizer/pp_calc.py b/atorch/utils/virtual_optimizer/pp_calc.py new file mode 100644 index 0000000..b3dc887 --- /dev/null +++ b/atorch/utils/virtual_optimizer/pp_calc.py @@ -0,0 +1,92 @@ +def is_valid_pipeline_parallel_combination( + num_layers, + pipeline_model_parallel_size, + decoder_first_pipeline_num_layers, + decoder_last_pipeline_num_layers, + decoder_first_virtual_pipeline_num_layers, + decoder_last_virtual_pipeline_num_layers, + num_virtual_stages_per_pipeline_rank, +): + """Check if a combination of parameters is valid according to the constraints.""" + if pipeline_model_parallel_size <= 2: + return False + + numerator = num_layers - decoder_first_pipeline_num_layers - decoder_last_pipeline_num_layers + denominator = pipeline_model_parallel_size - 2 + + if num_virtual_stages_per_pipeline_rank <= 1: + return False + + if decoder_first_pipeline_num_layers <= num_virtual_stages_per_pipeline_rank: + return False + + if decoder_last_pipeline_num_layers <= num_virtual_stages_per_pipeline_rank: + return False + + if numerator % denominator != 0: + return False + + num_layers_per_pp_stage = numerator // denominator + + if num_layers_per_pp_stage < decoder_first_pipeline_num_layers: + return False + + if num_layers_per_pp_stage < decoder_last_pipeline_num_layers: + return False + + if num_layers_per_pp_stage > 15: + return False + + if decoder_first_pipeline_num_layers > num_layers_per_pp_stage: + return False + + if decoder_last_pipeline_num_layers > num_layers_per_pp_stage: + return False + + if num_layers_per_pp_stage % num_virtual_stages_per_pipeline_rank != 0: + return False + + num_layers_per_vpp_stage = num_layers_per_pp_stage // num_virtual_stages_per_pipeline_rank + decoder_first_pp_rank_num_layers_without_first_vpp_stage = ( + decoder_first_pipeline_num_layers - decoder_first_virtual_pipeline_num_layers + ) + + if num_virtual_stages_per_pipeline_rank > 1: + if decoder_first_pp_rank_num_layers_without_first_vpp_stage % (num_virtual_stages_per_pipeline_rank - 1) != 0: + return False + decoder_first_pp_rank_not_first_vpp_stage_num_layers = ( + decoder_first_pp_rank_num_layers_without_first_vpp_stage // (num_virtual_stages_per_pipeline_rank - 1) + ) + + if decoder_first_pp_rank_not_first_vpp_stage_num_layers > num_layers_per_vpp_stage: + return False + if decoder_first_pp_rank_not_first_vpp_stage_num_layers < decoder_first_virtual_pipeline_num_layers: + return False + else: + if num_layers_per_vpp_stage > decoder_first_pp_rank_num_layers_without_first_vpp_stage: + return False + + if decoder_first_virtual_pipeline_num_layers > num_layers_per_vpp_stage: + return False + + if num_layers_per_vpp_stage > decoder_first_virtual_pipeline_num_layers + 2: + return False + + decoder_last_pp_rank_not_last_vpp_stage_num_layers = ( + decoder_last_pipeline_num_layers - decoder_last_virtual_pipeline_num_layers + ) + + if decoder_last_pp_rank_not_last_vpp_stage_num_layers % (num_virtual_stages_per_pipeline_rank - 1) != 0: + return False + + decoder_last_pp_rank_not_last_vpp_stage_num_layers = decoder_last_pp_rank_not_last_vpp_stage_num_layers // ( + num_virtual_stages_per_pipeline_rank - 1 + ) + + if decoder_last_pp_rank_not_last_vpp_stage_num_layers > num_layers_per_vpp_stage: + return False + + if num_layers_per_vpp_stage > decoder_last_pp_rank_not_last_vpp_stage_num_layers + 2: + return False + + return True diff --git a/examples/atorch_trainer_v2/baseline_megatron/pretrain_llama2_7b.sh b/examples/atorch_trainer_v2/pretrain/baseline_megatron/pretrain_llama2_7b.sh similarity index 100% rename from examples/atorch_trainer_v2/baseline_megatron/pretrain_llama2_7b.sh rename to examples/atorch_trainer_v2/pretrain/baseline_megatron/pretrain_llama2_7b.sh diff --git a/examples/atorch_trainer_v2/gpt2_config.yaml b/examples/atorch_trainer_v2/pretrain/gpt2_config.yaml similarity index 100% rename from examples/atorch_trainer_v2/gpt2_config.yaml rename to examples/atorch_trainer_v2/pretrain/gpt2_config.yaml diff --git a/examples/atorch_trainer_v2/run.sh b/examples/atorch_trainer_v2/pretrain/launch_pretrain.sh similarity index 82% rename from examples/atorch_trainer_v2/run.sh rename to examples/atorch_trainer_v2/pretrain/launch_pretrain.sh index dcb2dc7..c2dd4e5 100755 --- a/examples/atorch_trainer_v2/run.sh +++ b/examples/atorch_trainer_v2/pretrain/launch_pretrain.sh @@ -1,17 +1,23 @@ #!/bin/bash set -exo pipefail +# GPU +# pip install deprecated tensorboard h5py + +# PPU +# pip install scipy tensorboardX mdatasets nltk h5py + CONFIG_PATH=${1:-"$(dirname $0)/llama2_7b_config.yaml"} -MEGATRON_BRANCH=${MEGATRON_BRANCH:-core_r0.11.0} -MEGATRON_PATH=/tmp/Megatron-LM-${MEGATRON_BRANCH} +MEGATRON_TAG=${MEGATRON_TAG:-ant_release_0.11.0_v1.6.0} +MEGATRON_PATH=/tmp/Megatron-LM-${MEGATRON_TAG} if [ ! -d ${MEGATRON_PATH} ]; then pushd $(dirname ${MEGATRON_PATH}) - git clone -b ${MEGATRON_BRANCH} https://code.alipay.com/Arc/Megatron-LM.git $(basename ${MEGATRON_PATH}) + git clone -b ${MEGATRON_TAG} https://code.alipay.com/Arc/Megatron-LM.git $(basename ${MEGATRON_PATH}) popd fi -export PYTHONPATH="$(dirname $0)/../../":${MEGATRON_PATH}:$PYTHONPATH +export PYTHONPATH="$(dirname $0)/../../../":${MEGATRON_PATH}:$PYTHONPATH export CUDA_DEVICE_MAX_CONNECTIONS=1 export TIME_STAMP=$(date '+%Y%m%d-%H%M%S') diff --git a/examples/atorch_trainer_v2/llama2_7b_config.yaml b/examples/atorch_trainer_v2/pretrain/llama2_7b_config.yaml similarity index 86% rename from examples/atorch_trainer_v2/llama2_7b_config.yaml rename to examples/atorch_trainer_v2/pretrain/llama2_7b_config.yaml index dbaea01..6dfd292 100644 --- a/examples/atorch_trainer_v2/llama2_7b_config.yaml +++ b/examples/atorch_trainer_v2/pretrain/llama2_7b_config.yaml @@ -1,24 +1,23 @@ ## training args output_dir: !ENV ${OUTPUT_DIR} -overwrite_output_dir: True -# resume_from_checkpoint: !ENV ${OUTPUT_DIR} +# overwrite_output_dir: True +resume_from_checkpoint: !ENV ${OUTPUT_DIR} # flash_checkpoint: True do_train: True do_eval: True distributed_type: "megatron" num_train_epochs: 6 -block_size: 512 -per_device_train_batch_size: 2 -per_device_eval_batch_size: 2 +# per_device_train_batch_size: 2 +# per_device_eval_batch_size: 2 preprocessing_num_workers: 6 learning_rate: 2.0e-5 weight_decay: 0.0 warmup_ratio: 0.03 -seed: 42 +seed: 1403 max_grad_norm: 0 bf16: True -save_strategy: "samples" # "no" "epoch" "steps" "samples" -save_samples: 240 +save_strategy: "steps" # "no" "epoch" "steps" "samples" +# save_samples: 240 save_steps: 1000 save_total_limit: 3 evaluation_strategy: "steps" # "no" "epoch" "steps" "samples" @@ -32,14 +31,14 @@ tensorboard_dir: !ENV ${OUTPUT_DIR}/runs/${TIME_STAMP} dataloader_num_workers: 0 gradient_checkpointing: True -extra_save_frequency_in_epoch: - - 0.2 - - 0.4 +extra_save_frequency_in_epoch: [0.25, 0.5] ## data args data_path: &data_path /hetero_infer/jinshi.cl/code/Megatron-LM-core_r0.6.0/wikitext-2-raw-v1-llama2/llama2_tokenized_train_text_document tokenizer_model: &tokenizer_model /hetero_infer/jinshi.cl/code/tokenizers/llama_tokenizer/tokenizer.model +dynamic_save_config_path: /tmp/dynamic_save_config.json + ## profiling args # profiler_type: nv # profiler_file_path: !ENV ${OUTPUT_DIR}/profiler_output @@ -49,6 +48,8 @@ tokenizer_model: &tokenizer_model /hetero_infer/jinshi.cl/code/tokenizers/llama_ # profiler_schedule_repeat: 1 # profiler_schedule_skip_first: 20 +use_deterministic_algorithms: true + extra_configs: model_type_name: "llama" # no_save_optim: true @@ -79,8 +80,8 @@ extra_configs: pretraining_flag: False use_mcore_models: True transformer_impl: "transformer_engine" - micro_batch_size: 2 - global_batch_size: 16 + micro_batch_size: 1 + global_batch_size: 8 add_bias_linear: False bias_gelu_fusion: False recompute_activations: True @@ -92,6 +93,7 @@ extra_configs: sequence_parallel: True distributed_backend: "nccl" use_distributed_optimizer: True + overlap_grad_reduce: true enable_one_logger: False log_timers_to_tensorboard: True log_validation_ppl_to_tensorboard: True diff --git a/examples/atorch_trainer_v2/pretrain_atorch_trainer_megatron.py b/examples/atorch_trainer_v2/pretrain/pretrain_atorch_trainer_megatron.py similarity index 98% rename from examples/atorch_trainer_v2/pretrain_atorch_trainer_megatron.py rename to examples/atorch_trainer_v2/pretrain/pretrain_atorch_trainer_megatron.py index 011d1fb..0c479d6 100644 --- a/examples/atorch_trainer_v2/pretrain_atorch_trainer_megatron.py +++ b/examples/atorch_trainer_v2/pretrain/pretrain_atorch_trainer_megatron.py @@ -36,16 +36,6 @@ from transformers.utils import check_min_version from transformers.utils.versions import require_version -try: - import ant_patches -except ModuleNotFoundError as e: - print(e) - print( - "Can't import ant_patches, if you want to use megatron with version >= 'core_r0.9.0', " - "please use 'ant_core_r0.9.0' branch." - ) - ant_patches = None - from atorch.common.log_utils import default_logger as logger from atorch.trainer.args import AtorchTrainingArgs from atorch.trainer.atorch_trainer_v2 import AtorchTrainerV2 @@ -53,6 +43,8 @@ from atorch.utils.import_util import is_megatron_lm_available from atorch.utils.version import get_megatron_version, is_megatron_version_bigger_than +# from .instruction_dataset_utils import InstructionDataset + if is_megatron_lm_available(): import megatron.legacy.model from megatron.core import mpu @@ -65,7 +57,9 @@ ) from megatron.core.transformer.spec_utils import import_module from megatron.legacy.data.data_samplers import MegatronPretrainingRandomSampler, MegatronPretrainingSampler - from megatron.training import get_args, get_tokenizer, print_rank_0 + from megatron.training import get_args + from megatron.training import get_tokenizer as megatron_get_tokenizer + from megatron.training import print_rank_0 from megatron.training.arguments import core_transformer_config_from_args from megatron.training.utils import get_batch_on_this_cp_rank, get_batch_on_this_tp_rank from megatron.training.yaml_arguments import core_transformer_config_from_yaml @@ -193,6 +187,7 @@ class DataTrainingArguments: default=None, metadata={"help": "The configuration name of the dataset to use (via the datasets library)."}, ) + dataset_path: Optional[str] = field(default=None, metadata={"help": "A dir containing dataset with .arrow format."}) train_file: Optional[str] = field(default=None, metadata={"help": "The input training data file (a text file)."}) validation_file: Optional[str] = field( default=None, @@ -248,6 +243,8 @@ class DataTrainingArguments: metadata={"help": "Whether to keep line breaks when using TXT files or not."}, ) + drop_last: bool = field(default=False, metadata={"help": "Whether to drop last batch in dataloader."}) + def __post_init__(self): if self.streaming: require_version("datasets>=2.0.0", "The streaming feature requires `datasets>=2.0.0`") @@ -350,7 +347,7 @@ def is_dataset_built_on_rank(): ) and mpu.get_tensor_model_parallel_rank() == 0 def core_gpt_dataset_config_from_args(args): - tokenizer = get_tokenizer() + tokenizer = megatron_get_tokenizer() if is_megatron_version_bigger_than("0.10.0"): from megatron.training.utils import get_blend_and_blend_per_split @@ -771,7 +768,6 @@ def loss_postprocessing(self, monitor_dict): class CustomCallback(TrainerCallback): def on_log(self, args, state, control, logs=None, **kwargs): pass - # print_rank_last(f"---> callback: {logs}") def main(): @@ -826,10 +822,10 @@ def replacer(match): train_valid_test_datasets_provider.is_distributed = True training_args.extra_configs["custom_model_provider_function"] = model_provider + training_args.extra_configs["custom_train_step_class"] = GPTTrainStep training_args.extra_configs["custom_megatron_dataloaders_provider_function"] = partial( build_train_valid_test_data_iterators, train_valid_test_datasets_provider ) - training_args.extra_configs["custom_train_step_class"] = GPTTrainStep # Initialize our Trainer trainer = AtorchTrainerV2( diff --git a/examples/atorch_trainer_v2/sft/launch_sft.sh b/examples/atorch_trainer_v2/sft/launch_sft.sh new file mode 100755 index 0000000..9003403 --- /dev/null +++ b/examples/atorch_trainer_v2/sft/launch_sft.sh @@ -0,0 +1,69 @@ +#!/bin/bash +set -exo pipefail + +# GPU +# pip install deprecated tensorboard h5py + +# PPU +# pip install scipy tensorboardX mdatasets nltk h5py + +CONFIG_PATH=${1:-"$(dirname $0)/llama2_7b_config.yaml"} + +MEGATRON_TAG=${MEGATRON_TAG:-ant_release_0.11.0_v1.6.0} +MEGATRON_PATH=/tmp/Megatron-LM-${MEGATRON_TAG} +if [ ! -d ${MEGATRON_PATH} ]; then + pushd $(dirname ${MEGATRON_PATH}) + git clone -b ${MEGATRON_TAG} https://code.alipay.com/Arc/Megatron-LM.git $(basename ${MEGATRON_PATH}) + popd +fi + +export PYTHONPATH="$(dirname $0)/../../":${MEGATRON_PATH}:$PYTHONPATH +export CUDA_DEVICE_MAX_CONNECTIONS=1 + +export TIME_STAMP=$(date '+%Y%m%d-%H%M%S') +export OUTPUT_DIR=${OUTPUT_DIR:-"/tmp/llama2_atorch_trainer/"} +if [ ! -d ${OUTPUT_DIR} ]; then + mkdir -p ${OUTPUT_DIR} +fi + +NODE_NAME=${POD_NAME:-"master-0"} + +GPUS_PER_NODE=${GPUS_PER_NODE:-8} + +if [[ $POD_NAME =~ "edljob" ]]; then + WORLD_SIZE=${WORKER_NUM:-1} + NODE_RANK=${RANK:-0} + MASTER_ADDR=${MASTER_ADDR:-127.0.0.1} + RANDOM_PORT=$[$RANDOM + 20000] + MASTER_PORT=${MASTER_PORT:-$RANDOM_PORT} + GPU_NUM=$((${GPUS_PER_NODE}*${WORLD_SIZE})) + echo "---> from edl runtime, WORLD_SIZE: ${WORLD_SIZE}, NODE_RANK: ${NODE_RANK}" + LAUNCHER=" \ + python -m atorch.distributed.run --fault_tolerant --network-check \ + --max_restarts=1 \ + --nnode=$WORLD_SIZE \ + --nproc_per_node=$GPUS_PER_NODE \ + --rdzv_conf join_timeout=300 \ + " +else + WORLD_SIZE=${WORLD_SIZE:-1} + NODE_RANK=${RANK:-0} + MASTER_ADDR=${MASTER_ADDR:-127.0.0.1} + RANDOM_PORT=$[$RANDOM + 20000] + MASTER_PORT=${MASTER_PORT:-$RANDOM_PORT} + GPU_NUM=$((${GPUS_PER_NODE}*${WORLD_SIZE})) + echo "---> from pytorch runtime, WORLD_SIZE: ${WORLD_SIZE}, NODE_RANK: ${NODE_RANK}, MASTER_ADDR: ${MASTER_ADDR}, MASTER_PORT: ${MASTER_PORT}" + LAUNCHER=" \ + torchrun \ + --nproc_per_node ${GPUS_PER_NODE} \ + --nnodes ${WORLD_SIZE} \ + --node_rank ${NODE_RANK} \ + --master_addr ${MASTER_ADDR} \ + --master_port ${MASTER_PORT} \ + " +fi + +CMD="${LAUNCHER[@]} $(dirname $0)/sft_atorch_trainer_megatron.py ${CONFIG_PATH}" + +echo ${CMD} +${CMD} 2>&1 | tee ${OUTPUT_DIR}/node_${NODE_RANK}_${TIME_STAMP}.log diff --git a/examples/atorch_trainer_v2/sft/llama2_7b_config.yaml b/examples/atorch_trainer_v2/sft/llama2_7b_config.yaml new file mode 100644 index 0000000..48f828f --- /dev/null +++ b/examples/atorch_trainer_v2/sft/llama2_7b_config.yaml @@ -0,0 +1,119 @@ +## training args +finetune_type: sft +model_name_or_path: /tmp/llama2_atorch_trainer_pretrained +dataset_path: /dnn_training_sys/dataset/nlp/alpaca/alpaca_data_cleaned.json +block_size: 4096 + +output_dir: !ENV ${OUTPUT_DIR} +# overwrite_output_dir: True +resume_from_checkpoint: !ENV ${OUTPUT_DIR} +# flash_checkpoint: True +do_train: True +do_eval: True +distributed_type: "megatron" +num_train_epochs: 6 +# per_device_train_batch_size: 2 +# per_device_eval_batch_size: 2 +preprocessing_num_workers: 6 +learning_rate: 2.0e-5 +weight_decay: 0.0 +warmup_ratio: 0.03 +seed: 1403 +max_grad_norm: 0 +bf16: True +save_strategy: "steps" # "no" "epoch" "steps" "samples" +# save_samples: 240 +save_steps: 1000 +save_total_limit: 3 +evaluation_strategy: "steps" # "no" "epoch" "steps" "samples" +eval_steps: 2000 +logging_strategy: "steps" # "no" "epoch" "steps" "samples" +logging_steps: 1 +logging_nan_inf_filter: False +# log_params_std: True +# log_grad_diff_for_debug: True +tensorboard_dir: !ENV ${OUTPUT_DIR}/runs/${TIME_STAMP} +dataloader_num_workers: 0 +gradient_checkpointing: True + +extra_save_frequency_in_epoch: [10, 50] + +## data args +data_path: &data_path /hetero_infer/jinshi.cl/code/Megatron-LM-core_r0.6.0/wikitext-2-raw-v1-llama2/llama2_tokenized_train_text_document +tokenizer_model: &tokenizer_model /hetero_infer/jinshi.cl/code/tokenizers/llama_tokenizer/tokenizer.model + +dynamic_save_config_path: /tmp/dynamic_save_config.json + +## profiling args +# profiler_type: nv +# profiler_file_path: !ENV ${OUTPUT_DIR}/profiler_output +# profiler_schedule_wait: 1 +# profiler_schedule_warmup: 1 +# profiler_schedule_active: 1 +# profiler_schedule_repeat: 1 +# profiler_schedule_skip_first: 20 + +use_deterministic_algorithms: true + +extra_configs: + model_type_name: "llama" + # no_save_optim: true + num_layers: 32 + hidden_size: 4096 + ffn_hidden_size: 11008 + num_attention_heads: 32 + group_query_attention: True + num_query_groups: 32 + max_position_embeddings: 4096 + position_embedding_type: "rope" + make_vocab_size_divisible_by: 1 + norm_epsilon: 1.0e-5 + normalization: "RMSNorm" + swiglu: True + untie_embeddings_and_output_weights: True + use_flash_attn: True + tokenizer_type: "Llama2Tokenizer" + tokenizer_model: *tokenizer_model + optimizer: "adam" + attention_dropout: 0.0 + hidden_dropout: 0.0 + weight_decay: 1.0e-1 + clip_grad: 1.0 + adam_beta1: 0.9 + adam_beta2: 0.95 + adam_eps: 1.0e-8 + pretraining_flag: False + use_mcore_models: True + transformer_impl: "transformer_engine" + micro_batch_size: 1 + global_batch_size: 8 + add_bias_linear: False + bias_gelu_fusion: False + recompute_activations: True + recompute_granularity: "selective" + # train_iters: 10000 + eval_iters: 50 + tensor_model_parallel_size: 1 + pipeline_model_parallel_size: 4 + num_virtual_stages_per_pipeline_rank: 4 + sequence_parallel: True + distributed_backend: "nccl" + use_distributed_optimizer: True + overlap_grad_reduce: true + enable_one_logger: False + log_timers_to_tensorboard: True + log_validation_ppl_to_tensorboard: True + log_memory_to_tensorboard: True + log_throughput: True + log_params_norm: True + seed: 1403 + init_method_std: 0.02 + lr: 3.0e-5 + min_lr: 3.0e-6 + lr_ecay_style: "cosine" + lr_warmup_fraction: 0.1 + data_path: [*data_path] + split: "949,50,1" + data_cache_path: !ENV ${OUTPUT_DIR}/data_cache + seq_length: 4096 + num_workers: 0 diff --git a/examples/atorch_trainer_v2/sft/sft_atorch_trainer_megatron.py b/examples/atorch_trainer_v2/sft/sft_atorch_trainer_megatron.py new file mode 100644 index 0000000..351b990 --- /dev/null +++ b/examples/atorch_trainer_v2/sft/sft_atorch_trainer_megatron.py @@ -0,0 +1,915 @@ +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2020 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" +Fine-tuning the library models for causal language modeling (GPT, GPT-2, CTRL, ...) on a text file or a dataset. + +Here is the full list of checkpoints on the hub that can be fine-tuned by this script: +https://huggingface.co/models?filter=text-generation +""" +# You can also adapt this script on your own causal language modeling task. Pointers for this are left as comments. + +import copy +import json +import logging +import os +import re +import sys +from dataclasses import dataclass, field +from functools import partial +from typing import Optional, Union + +import torch +import yaml # type: ignore[import] +from torch.utils.data import DataLoader, Dataset +from torch.utils.data.distributed import DistributedSampler +from transformers import MODEL_FOR_CAUSAL_LM_MAPPING, AutoTokenizer, HfArgumentParser, TrainerCallback +from transformers.utils import check_min_version +from transformers.utils.versions import require_version + +from atorch.common.log_utils import default_logger as logger +from atorch.trainer.args import AtorchTrainingArgs +from atorch.trainer.atorch_trainer_v2 import AtorchTrainerV2 +from atorch.trainer.megatron import MegatronTrainStep +from atorch.utils.import_util import is_megatron_lm_available + +if is_megatron_lm_available(): + import megatron.legacy.model + from megatron.core import mpu + from megatron.core.models.gpt import GPTModel + from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_local_spec, + get_gpt_layer_with_transformer_engine_spec, + ) + from megatron.core.transformer.spec_utils import import_module + from megatron.training import get_args + from megatron.training import get_tokenizer as megatron_get_tokenizer + from megatron.training import print_rank_0 + from megatron.training.arguments import core_transformer_config_from_args + from megatron.training.yaml_arguments import core_transformer_config_from_yaml + +# Will error if the minimal version of Transformers is not installed. Remove at your own risks. +check_min_version("4.31.0") + +require_version( + "datasets>=2.14.0", + "To fix: pip install -r examples/pytorch/language-modeling/requirements.txt", +) + + +MODEL_CONFIG_CLASSES = list(MODEL_FOR_CAUSAL_LM_MAPPING.keys()) +MODEL_TYPES = tuple(conf.model_type for conf in MODEL_CONFIG_CLASSES) + + +@dataclass +class ModelArguments: + """ + Arguments pertaining to which model/config/tokenizer we are going to fine-tune, or train from scratch. + """ + + model_name_or_path: Optional[str] = field( + default=None, + metadata={ + "help": ( + "The model checkpoint for weights initialization. Don't set if you want to train a model from scratch." + ) + }, + ) + vocab_file: Optional[str] = field(default=None, metadata={"help": "The vocab file (a json file)."}) + merge_file: Optional[str] = field(default=None, metadata={"help": "The merge file (a json file)."}) + model_type: Optional[str] = field( + default=None, + metadata={"help": "If training from scratch, pass a model type from the list: " + ", ".join(MODEL_TYPES)}, + ) + config_overrides: Optional[str] = field( + default=None, + metadata={ + "help": ( + "Override some existing default config settings when a model is trained from scratch. Example: " + "n_embd=10,resid_pdrop=0.2,scale_attn_weights=false,summary_type=cls_index" + ) + }, + ) + config_name: Optional[str] = field( + default=None, + metadata={"help": "Pretrained config name or path if not the same as model_name"}, + ) + tokenizer_name: Optional[str] = field( + default=None, + metadata={"help": "Pretrained tokenizer name or path if not the same as model_name"}, + ) + tokenizer_model: Optional[str] = field(default=None, metadata={"help": "The path to tokenizer model."}) + cache_dir: Optional[str] = field( + default=None, + metadata={"help": "Where do you want to store the pretrained models downloaded from huggingface.co"}, + ) + use_fast_tokenizer: bool = field( + default=True, + metadata={"help": "Whether to use one of the fast tokenizer (backed by the tokenizers library) or not."}, + ) + model_revision: str = field( + default="main", + metadata={"help": "The specific model version to use (can be a branch name, tag name or commit id)."}, + ) + token: str = field( + default=None, + metadata={ + "help": ( + "The token to use as HTTP bearer authorization for remote files. If not specified, will use the token " + "generated when running `huggingface-cli login` (stored in `~/.huggingface`)." + ) + }, + ) + trust_remote_code: bool = field( + default=False, + metadata={ + "help": ( + "Whether to trust the execution of code from datasets/models defined on the Hub." + " This option should only be set to `True` for repositories you trust and in which you have read the" + " code, as it will execute code present on the Hub on your local machine." + ) + }, + ) + torch_dtype: Optional[str] = field( + default=None, + metadata={ + "help": ( + "Override the default `torch.dtype` and load the model under this dtype. If `auto` is passed, the " + "dtype will be automatically derived from the model's weights." + ), + "choices": ["auto", "bfloat16", "float16", "float32"], + }, + ) + low_cpu_mem_usage: bool = field( + default=False, + metadata={ + "help": ( + "It is an option to create the model as an empty shell, then only materialize its parameters when the " + "pretrained weights are loaded. set True will benefit LLM loading time and RAM consumption." + ) + }, + ) + + def __post_init__(self): + if self.config_overrides is not None and (self.config_name is not None or self.model_name_or_path is not None): + raise ValueError( + "--config_overrides can't be used in combination with --config_name or --model_name_or_path" + ) + + +@dataclass +class DataTrainingArguments: + """ + Arguments pertaining to what data we are going to input our model for training and eval. + """ + + dataset_name: Optional[str] = field( + default=None, + metadata={"help": "The name of the dataset to use (via the datasets library)."}, + ) + dataset_config_name: Optional[str] = field( + default=None, + metadata={"help": "The configuration name of the dataset to use (via the datasets library)."}, + ) + dataset_path: Optional[str] = field(default=None, metadata={"help": "A dir containing dataset with .arrow format."}) + train_file: Optional[str] = field(default=None, metadata={"help": "The input training data file (a text file)."}) + validation_file: Optional[str] = field( + default=None, + metadata={"help": "An optional input evaluation data file to evaluate the perplexity on (a text file)."}, + ) + data_path: Optional[str] = field( + default=None, + metadata={"help": "The input training data path (preprocessed Megatron data)."}, + ) + max_train_samples: Optional[int] = field( + default=None, + metadata={ + "help": ( + "For debugging purposes or quicker training, truncate the number of training examples to this " + "value if set." + ) + }, + ) + max_eval_samples: Optional[int] = field( + default=None, + metadata={ + "help": ( + "For debugging purposes or quicker training, truncate the number of evaluation examples to this " + "value if set." + ) + }, + ) + streaming: bool = field(default=False, metadata={"help": "Enable streaming mode"}) + block_size: Optional[int] = field( + default=None, + metadata={ + "help": ( + "Optional input sequence length after tokenization. " + "The training dataset will be truncated in block of this size for training. " + "Default to the model max input length for single sentence inputs (take into account special tokens)." + ) + }, + ) + overwrite_cache: bool = field( + default=False, + metadata={"help": "Overwrite the cached training and evaluation sets"}, + ) + validation_split_percentage: Optional[int] = field( + default=5, + metadata={"help": "The percentage of the train set used as validation set in case there's no validation split"}, + ) + preprocessing_num_workers: Optional[int] = field( + default=None, + metadata={"help": "The number of processes to use for the preprocessing."}, + ) + keep_linebreaks: bool = field( + default=True, + metadata={"help": "Whether to keep line breaks when using TXT files or not."}, + ) + + drop_last: bool = field(default=False, metadata={"help": "Whether to drop last batch in dataloader."}) + + def __post_init__(self): + if self.streaming: + require_version("datasets>=2.0.0", "The streaming feature requires `datasets>=2.0.0`") + + if ( + self.dataset_name is None + and self.train_file is None + and self.validation_file is None + and self.data_path is None + ): + raise ValueError("Need either a dataset name or a training/validation file.") + else: + if self.train_file is not None: + extension = self.train_file.split(".")[-1] + assert extension in [ + "csv", + "json", + "txt", + ], "`train_file` should be a csv, a json or a txt file." + if self.validation_file is not None: + extension = self.validation_file.split(".")[-1] + assert extension in [ + "csv", + "json", + "txt", + ], "`validation_file` should be a csv, a json or a txt file." + + +def model_provider(pre_process=True, post_process=True) -> Union[GPTModel, megatron.legacy.model.GPTModel]: + """Builds the model. + + If you set the use_mcore_models to True, it will return the mcore GPT model and if not the legacy GPT model. + + Args: + pre_process (bool, optional): Set to true if you need to compute embedings. Defaults to True. + post_process (bool, optional): Set to true if you need to want to compute output logits/loss. Defaults to True. + + + Returns: + Union[GPTModel, megatron.legacy.model.GPTModel]: The returned model + """ + args = get_args() + use_te = args.transformer_impl == "transformer_engine" + + print_rank_0("building GPT model ...") + # Experimental loading arguments from yaml + if args.yaml_cfg is not None: + config = core_transformer_config_from_yaml(args, "language_model") + else: + config = core_transformer_config_from_args(args) + + if args.use_mcore_models: + if args.spec is not None: + transformer_layer_spec = import_module(args.spec) + else: + if use_te: + transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( + args.num_experts, args.moe_grouped_gemm + ) + else: + transformer_layer_spec = get_gpt_layer_local_spec(args.num_experts, args.moe_grouped_gemm) + + model = GPTModel( + config=config, + transformer_layer_spec=transformer_layer_spec, + vocab_size=args.padded_vocab_size, + max_sequence_length=args.max_position_embeddings, + pre_process=pre_process, + post_process=post_process, + fp16_lm_cross_entropy=args.fp16_lm_cross_entropy, + parallel_output=True, + share_embeddings_and_output_weights=not args.untie_embeddings_and_output_weights, + position_embedding_type=args.position_embedding_type, + rotary_percent=args.rotary_percent, + ) + else: + assert args.context_parallel_size == 1, "Context parallelism is only supported with Megatron Core!" + + model = megatron.legacy.model.GPTModel( + config, + num_tokentypes=0, + parallel_output=True, + pre_process=pre_process, + post_process=post_process, + ) + + return model + + +def is_dataset_built_on_rank(): + return (mpu.is_pipeline_first_stage() or mpu.is_pipeline_last_stage()) and mpu.get_tensor_model_parallel_rank() == 0 + + +class GPTTrainStep(MegatronTrainStep): + """ + GPT train step + + Args: + args (`argparse.Namespace`): Megatron-LM arguments. + """ + + def __init__(self, args, **kwargs): + super().__init__() + if not args.model_return_dict: + self.model_output_class = None + else: + from transformers.modeling_outputs import CausalLMOutputWithCrossAttentions + + self.model_output_class = CausalLMOutputWithCrossAttentions + + ################################## + # Just for testing spike loss + self.last_loss = None + # Just for testing spike loss + ################################## + + def get_batch_func(self, **kwargs): + def get_batch(data_iterator): + """Generate a batch.""" + + # TODO: this is pretty hacky, find a better way + if (not mpu.is_pipeline_first_stage()) and (not mpu.is_pipeline_last_stage()): + return None, None, None, None, None + + # get batches based on the TP rank you are on + # batch = get_batch_on_this_tp_rank(data_iterator) + + # # slice batch along sequence dimension for context parallelism + # batch = get_batch_on_this_cp_rank(batch) + + batch = sft_get_batch_on_this_tp_rank(data_iterator) + + return batch.values() + + return get_batch + + def get_loss_func(self, **kwargs): + def loss_func(loss_mask: torch.Tensor, output_tensor: torch.Tensor): + """Loss function. + + Args: + loss_mask (torch.Tensor): Used to mask out some portions of the loss + output_tensor (torch.Tensor): The tensor with the losses + + Returns: + the loss scalar for this micro-batch + the number of non-padded tokens in this microbatch + a dict containing reporting metrics on the loss and number of tokens across + the data parallel ranks + """ + args = get_args() + + losses = output_tensor.float() + loss_mask = loss_mask.view(-1).float() + total_tokens = loss_mask.sum() + loss = torch.cat([torch.sum(losses.view(-1) * loss_mask).view(1), total_tokens.view(1)]) + + if args.context_parallel_size > 1: + torch.distributed.all_reduce(loss, group=mpu.get_context_parallel_group()) + + # Reduce loss for logging. + reporting_loss = loss.clone().detach() + torch.distributed.all_reduce(reporting_loss, group=mpu.get_data_parallel_group()) + + local_num_tokens = loss[1].clone().detach().to(torch.int) + return ( + loss[0] * args.context_parallel_size, + local_num_tokens, + {"lm loss": (reporting_loss[0], reporting_loss[1])}, + ) + + return loss_func + + def get_forward_step_func(self, **kwargs): + def forward_step(data_iterator, model: GPTModel): + """Forward training step. + + Args: + data_iterator : Input data iterator + model (GPTModel): The GPT Model + """ + # Get the batch. + tokens, labels, loss_mask, attention_mask, position_ids = self.get_batch_func()(data_iterator) + output_tensor = model(tokens, position_ids, attention_mask, labels=labels) + + return output_tensor, partial(self.get_loss_func(), loss_mask) + + return forward_step + + def loss_postprocessing(self, monitor_dict): + """ + Loss postprocessing. Average losses across all micro-batches. + + Args: + monitor_dict: a dict to hold all related monitored matrix + In train process: {losses_reduced: [], total_grad_norm: float or None} + losses_reduced: (List[torch.Tensor]): + A list of losses whose length equals to the number of microbatches, `global_batch_size/data_parallel_size/micro_batch_size`. # noqa E501 + total_grad_norm (Only in training process): + total grad_norm (for all params). + In eval or test process: {losses_reduced: []} + losses_reduced: (List[torch.Tensor]) + Returns: + A dict: + In train process: return {"loss_to_log": Dict, "spike_loss_ratio": float or None} + The first one is a train loss dict to log, and the second one is a ratio if the spike loss occurs. + In eval or test process: return {"loss_to_log": Dict} + """ + + args = get_args() + + assert "losses_reduced" in monitor_dict + + losses_reduced = monitor_dict.get("losses_reduced", None) + + # args.forward_mode is a internal intermediate variable to indicate current forward mode. + # args.forward_mode belongs to ["train", "eval", "test"] + if args.forward_mode == "train": + total_grad_norm = monitor_dict.get("total_grad_norm", None) # noqa: F841 + res = {"loss_to_log": {}, "spike_loss_ratio": None} + else: + res = {"loss_to_log": {}} + + if mpu.is_pipeline_last_stage(ignore_virtual=True): + # Average loss across microbatches. + loss_reduced = {} + for key in losses_reduced[0].keys(): + numerator = 0 + denominator = 0 + for x in losses_reduced: + val = x[key] + # there is one dict per microbatch. in new reporting, we average + # over the total number of tokens across the global batch. + if isinstance(val, tuple) or isinstance(val, list): + numerator += val[0] + denominator += val[1] + else: + # legacy behavior. we average over the number of microbatches, + # and so the denominator is 1. + numerator += val + denominator += 1 + loss_reduced[key] = numerator / denominator + ################################## + # Just for testing spike loss + ratio = None + if ( + self.last_loss is not None + and loss_reduced["lm loss"] > self.last_loss + and torch.abs(loss_reduced["lm loss"] - self.last_loss) > 1 + ): + ratio = 0.8 + spike_loss = loss_reduced["lm loss"] * ratio + loss_reduced.update(spike_loss=spike_loss) + logger.info( + f' current loss {loss_reduced["lm loss"]}, last loss: {self.last_loss}, spike_loss: {spike_loss}' + ) + self.last_loss = loss_reduced["lm loss"] + res["loss_to_log"] = loss_reduced + res["spike_loss_ratio"] = ratio + # Just for testing spike loss + ################################## + return res + + +class CustomCallback(TrainerCallback): + def on_log(self, args, state, control, logs=None, **kwargs): + pass + + +class MegatronDistributedSampler(DistributedSampler): + def __init__(self, dataset, num_replicas=None, rank=None, shuffle=True, seed=0, drop_last=False): + + self.dp_world_size = mpu.get_data_parallel_world_size() + self.rank = mpu.get_data_parallel_rank() + super().__init__( + dataset, num_replicas=self.dp_world_size, rank=self.rank, shuffle=shuffle, seed=seed, drop_last=drop_last + ) + + +PROMPT_DICT = { + "prompt_input": ( + "Below is an instruction that describes a task, paired with an input that provides further context. " + "Write a response that appropriately completes the request.\n\n" + "### Instruction:\n{instruction}\n\n### Input:\n{input}\n\n### Response:" + ), + "prompt_no_input": ( + "Below is an instruction that describes a task. " + "Write a response that appropriately completes the request.\n\n" + "### Instruction:\n{instruction}\n\n### Response:" + ), +} + + +# copy from llama-recipes +# https://github.com/facebookresearch/llama-recipes/blob/405255c/ft_datasets/alpaca_dataset.py +class InstructionDataset(Dataset): + def __init__(self, dataset_path, tokenizer, partition="train", max_words=30): + self.ann = json.load(open(dataset_path)) + if partition == "train": + self.ann = self.ann + else: + self.ann = self.ann[:200] + self.max_words = max_words + self.tokenizer = tokenizer + + def __len__(self): + return len(self.ann) + + def __getitem__(self, index): + + args = get_args() + + logger.info(f"global step {args.global_step} rank {args.rank}, index {index}") + + IGNORE_INDEX = -100 # The default setting in CrossEntropyLoss + ann = self.ann[index] + if ann.get("input", "") == "": + prompt = PROMPT_DICT["prompt_no_input"].format_map(ann) + else: + prompt = PROMPT_DICT["prompt_input"].format_map(ann) + example = prompt + ann["output"] + + ############ Original InstructionDataset + # prompt = torch.tensor(self.tokenizer.encode(prompt), dtype=torch.int64) + # example = self.tokenizer.encode(example) + # example.append(self.tokenizer.eos_token_id) + + ############ Add by jinshi.cl, for Megatron ############ + prompt = torch.tensor(self.tokenizer.tokenize(prompt), dtype=torch.int64) + example = self.tokenizer.tokenize(example) + example.append(self.tokenizer.eos_id) + + example = torch.tensor(example, dtype=torch.int64) + padding = self.max_words - example.shape[0] + if padding > 0: + example = torch.cat((example, torch.zeros(padding, dtype=torch.int64) - 1)) + elif padding < 0: + example = example[: self.max_words] + labels = copy.deepcopy(example) + labels[: len(prompt)] = -1 + example_mask = example.ge(0) + label_mask = labels.ge(0) + example[~example_mask] = 0 + labels[~label_mask] = IGNORE_INDEX + example_mask = example_mask.float() + label_mask = label_mask.float() + + ############ Original InstructionDataset + # return { + # "input_ids": example, + # "labels": labels, + # "attention_mask": example_mask, + # } + + ############ Add by jinshi.cl, for Megatron ############ + return { + "tokens": example, + "labels": labels, + "loss_mask": label_mask, + } + + +def get_tokenizer(model_args): + tokenizer_kwargs = { + "cache_dir": model_args.cache_dir, + "use_fast": model_args.use_fast_tokenizer, + "revision": model_args.model_revision, + "use_auth_token": True if model_args.use_auth_token else None, + } + if model_args.tokenizer_name: + tokenizer = AutoTokenizer.from_pretrained(model_args.tokenizer_name, trust_remote_code=True, **tokenizer_kwargs) + elif model_args.model_name_or_path: + tokenizer = AutoTokenizer.from_pretrained( + model_args.model_name_or_path, trust_remote_code=True, **tokenizer_kwargs + ) + else: + raise ValueError( + "You are instantiating a new tokenizer from scratch. This is not supported by this script." + "You can do it from another script, save it, and load it from here, using --tokenizer_name." + ) + return tokenizer + + +def build_train_valid_test_data_iterators_for_sft(model_args, data_args, training_args: AtorchTrainingArgs): + from megatron.training.global_vars import get_args + + args = get_args() + + # tokenizer = get_tokenizer(model_args) + tokenizer = megatron_get_tokenizer() + + if not is_dataset_built_on_rank(): + print(f"dataset: rank {torch.distributed.get_rank()}, build empty dataloader.") + return None, None, None + else: + train_dataset = InstructionDataset( + data_args.dataset_path, + tokenizer, + partition="train", + max_words=data_args.block_size, + ) + + eval_dataset = InstructionDataset( + data_args.dataset_path, + tokenizer, + partition="eval", + max_words=data_args.block_size, + ) + + batch_size = args.micro_batch_size + + train_dataloader = DataLoader( + train_dataset, + sampler=MegatronDistributedSampler( + train_dataset, shuffle=True, seed=training_args.seed, drop_last=data_args.drop_last + ), + batch_size=batch_size, + pin_memory=True, + drop_last=data_args.drop_last, + ) + + eval_dataloader = DataLoader( + eval_dataset, + sampler=MegatronDistributedSampler(eval_dataset), + batch_size=batch_size, + pin_memory=True, + drop_last=data_args.drop_last, + ) + + print( + f"[Rank {torch.distributed.get_rank()}] train_dataloader_type {type(train_dataloader)} train_dataset " + f"{len(train_dataset)} train_dataloader {len(train_dataloader)} eval_dataset {len(eval_dataset)}", + flush=True, + ) + return train_dataloader, eval_dataloader, None + + +def sft_get_batch_on_this_tp_rank(data_iterator): + + args = get_args() + + def _broadcast(item): + if item is not None: + torch.distributed.broadcast( + item, mpu.get_tensor_model_parallel_src_rank(), group=mpu.get_tensor_model_parallel_group() + ) + + if mpu.get_tensor_model_parallel_rank() == 0: + + if data_iterator is not None: + data = next(data_iterator) + else: + data = None + + batch = { + "tokens": data["tokens"].cuda(non_blocking=True), + "labels": data["labels"].cuda(non_blocking=True), + "loss_mask": data["loss_mask"].cuda(non_blocking=True), + "attention_mask": None, + "position_ids": None + # 'attention_mask': None if "attention_mask" not in data else data["attention_mask"].cuda(non_blocking = True), # noqa: E501 + # 'position_ids': data["position_ids"].cuda(non_blocking = True) + } + + if args.pipeline_model_parallel_size == 1: + _broadcast(batch["tokens"]) + _broadcast(batch["labels"]) + _broadcast(batch["loss_mask"]) + # _broadcast(batch['attention_mask']) + # _broadcast(batch['position_ids']) + + elif mpu.is_pipeline_first_stage(): + _broadcast(batch["tokens"]) + # _broadcast(batch['attention_mask']) + # _broadcast(batch['position_ids']) + + elif mpu.is_pipeline_last_stage(): + _broadcast(batch["labels"]) + _broadcast(batch["loss_mask"]) + # _broadcast(batch['attention_mask']) + + else: + + tokens = torch.empty( + (args.micro_batch_size, args.seq_length), dtype=torch.int64, device=torch.cuda.current_device() + ) + labels = torch.empty( + (args.micro_batch_size, args.seq_length), dtype=torch.int64, device=torch.cuda.current_device() + ) + loss_mask = torch.empty( + (args.micro_batch_size, args.seq_length), dtype=torch.float32, device=torch.cuda.current_device() + ) + # if args.create_attention_mask_in_dataloader: + # attention_mask=torch.empty( + # (args.micro_batch_size,1,args.seq_length,args.seq_length), dtype = torch.bool , device = torch.cuda.current_device() # noqa: E501 + # ) + # else: + # attention_mask=None + # position_ids=torch.empty((args.micro_batch_size,args.seq_length), dtype = torch.int64 , device = torch.cuda.current_device()) # noqa: E501 + + if args.pipeline_model_parallel_size == 1: + _broadcast(tokens) + _broadcast(labels) + _broadcast(loss_mask) + # _broadcast(attention_mask) + # _broadcast(position_ids) + + elif mpu.is_pipeline_first_stage(): + labels = None + loss_mask = None + + _broadcast(tokens) + # _broadcast(attention_mask) + # _broadcast(position_ids) + + elif mpu.is_pipeline_last_stage(): + tokens = None + # position_ids = None + + _broadcast(labels) + _broadcast(loss_mask) + # _broadcast(attention_mask) + + batch = { + "tokens": tokens, + "labels": labels, + "loss_mask": loss_mask, + "attention_mask": None, + "position_ids": None + # 'attention_mask': attention_mask, + # 'position_ids': position_ids + } + + return batch + + +def main(): + # See all possible arguments in src/transformers/training_args.py + # or by passing the --help flag to this script. + # We now keep distinct sets of args, for a cleaner separation of concerns. + + parser = HfArgumentParser((ModelArguments, DataTrainingArguments, AtorchTrainingArgs)) + if len(sys.argv) == 2: + # If we pass only one argument to the script and it's the path to a json or yaml file, + # let's parse it to get our arguments. + if sys.argv[1].endswith(".json"): + model_args, data_args, training_args = parser.parse_json_file(json_file=os.path.abspath(sys.argv[1])) + elif sys.argv[1].endswith(".yaml") or sys.argv[1].endswith(".yml"): + + def resolve_env_vars(loader, node): + value = node.value + if not isinstance(value, str): + return value + + # parse ${VAR:-default_value} + pattern = re.compile(r"\$\{([^:}]+)(?::-(.+?))?\}") + + def replacer(match): + var_name = match.group(1) # 环境变量名 + default_value = match.group(2) # 默认值 + # 返回环境变量值或默认值 + print(var_name, os.getenv(var_name, default_value)) + return os.getenv(var_name, default_value) + + return pattern.sub(replacer, value) + + yaml.SafeLoader.add_constructor("!ENV", resolve_env_vars) + + model_args, data_args, training_args = parser.parse_yaml_file(yaml_file=os.path.abspath(sys.argv[1])) + else: + model_args, data_args, training_args = parser.parse_args_into_dataclasses() + + # Sending telemetry. Tracking the example usage helps us better allocate resources to maintain them. The + # information sent is the one passed as arguments along with your Python/PyTorch versions. + # send_example_telemetry("run_clm", model_args, data_args) + + # Setup logging + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + handlers=[logging.StreamHandler(sys.stdout)], + ) + + logger.info(f"Training/evaluation parameters {training_args}") + + assert training_args.finetune_type is not None, "finetune_type should be set!" + + training_args.extra_configs["custom_model_provider_function"] = model_provider + training_args.extra_configs["custom_train_step_class"] = GPTTrainStep + training_args.extra_configs["custom_megatron_dataloaders_provider_function"] = partial( + build_train_valid_test_data_iterators_for_sft, model_args, data_args, training_args + ) + + if model_args.model_name_or_path is not None: + tracker_filename = os.path.join(model_args.model_name_or_path, "latest_checkpointed_iteration.txt") + assert os.path.exists( + tracker_filename + ), f"latest_checkpointed_iteration.txt should be in {model_args.model_name_or_path}." + with open(tracker_filename, "r") as f: + metastring = f.read().strip() + assert ( + metastring == "release" + ), f"The content of {tracker_filename} should be 'release', but got {metastring}." + pretrained_model_dir = os.path.join(model_args.model_name_or_path, "release") + assert os.path.isdir(pretrained_model_dir), f"pretrained model dir {pretrained_model_dir} does not exist!" + + # If resume from checkpoint, training_args.resume_from_checkpoint is equal to training_args.output_dir + if training_args.resume_from_checkpoint is None: + training_args.resume_from_checkpoint = model_args.model_name_or_path + logger.info(f"Train from pretrained model {training_args.resume_from_checkpoint}.") + else: + assert training_args.resume_from_checkpoint == training_args.output_dir, ( + "training_args.resume_from_checkpoint should be equal to training_args.output_dir, " + f"but got {training_args.resume_from_checkpoint} and {training_args.output_dir}." + ) + tracker_filename = os.path.join(training_args.resume_from_checkpoint, "latest_checkpointed_iteration.txt") + if os.path.exists(tracker_filename): + with open(tracker_filename, "r") as f: + metastring = f.read().strip() + try: + iteration = int(metastring) + assert iteration > 0 + except ValueError: + logger.error(f"iteration value in {tracker_filename} should be greater than 0!") + raise + resumed_ckpt = os.path.join(training_args.resume_from_checkpoint, f"iter_{iteration:07d}") + logger.info(f"Train from resumed model {resumed_ckpt}.") + else: + training_args.resume_from_checkpoint = model_args.model_name_or_path + logger.info(f"Train from pretrained model {training_args.resume_from_checkpoint}.") + + # Initialize our Trainer + trainer = AtorchTrainerV2( + args=training_args, + callbacks=[CustomCallback()], + ) + + # Training + if training_args.do_train: + train_result = trainer.train() + + metrics = train_result.metrics # noqa F401 + + # max_train_samples = ( + # data_args.max_train_samples if data_args.max_train_samples is not None else len(train_dataset) + # ) + # metrics["train_samples"] = min(max_train_samples, len(train_dataset)) + + # trainer.log_metrics("train", metrics) + # trainer.save_metrics("train", metrics) + # trainer.save_state() + + # Evaluation + if training_args.do_eval: + logger.info("*** Evaluate ***") + + metrics = trainer.evaluate() # noqa: F841 + + # max_eval_samples = data_args.max_eval_samples if data_args.max_eval_samples is not None else len(eval_dataset) + # metrics["eval_samples"] = min(max_eval_samples, len(eval_dataset)) + # try: + # perplexity = math.exp(metrics["eval_loss"]) + # except OverflowError: + # perplexity = float("inf") + # metrics["perplexity"] = perplexity + + # trainer.log_metrics("eval", metrics) + # trainer.save_metrics("eval", metrics) + + +if __name__ == "__main__": + main() diff --git a/examples/moe/moe_modules.py b/examples/moe/moe_modules.py index 48ccfc3..9ea7fed 100644 --- a/examples/moe/moe_modules.py +++ b/examples/moe/moe_modules.py @@ -3,11 +3,19 @@ import torch import torch.nn.functional as F +import transformers +from packaging.version import Version from torch import nn from torch.utils.data import DataLoader, Dataset from torch.utils.data.distributed import DistributedSampler from transformers.models.llama import modeling_llama -from transformers.models.llama.modeling_llama import LlamaDecoderLayer, LlamaForCausalLM, LlamaModel, LlamaRMSNorm +from transformers.models.llama.modeling_llama import ( + LlamaDecoderLayer, + LlamaForCausalLM, + LlamaModel, + LlamaRMSNorm, + LlamaRotaryEmbedding, +) from atorch.common.util_func import data_to_device from atorch.distributed.distributed import get_device_mesh @@ -219,7 +227,11 @@ def __init__(self, config, layer_num=None, pre_process=True, post_process=True, self.layer_num = config.num_hidden_layers else: self.layer_num = layer_num - + if Version(transformers.__version__) >= Version("4.48.0"): + # 4.48.0 之后,传入 decoder_layer 的 position_embeddings 不能为 None。 + self.rotary_emb = LlamaRotaryEmbedding(config=config) + else: + self.rotary_emb = None self.layers = torch.nn.ModuleDict() for layer_idx_cur_stage in range(self.layer_num): @@ -252,10 +264,15 @@ def forward(self, input_ids): position_ids = torch.arange(0, hidden_states.shape[1], device=hidden_states.device) position_ids = position_ids.unsqueeze(0) + position_embeddings = self.rotary_emb and self.rotary_emb(hidden_states, position_ids) for decoder_layer in self.layers.values(): - hidden_states = decoder_layer(hidden_states, position_ids=position_ids)[0] + hidden_states = decoder_layer( + hidden_states, + position_ids=position_ids, + position_embeddings=position_embeddings, + )[0] - if self.norm is not None: + if self.norm is not None and self.lm_head is not None: hidden_states = self.norm(hidden_states) logits = self.lm_head(hidden_states) return logits From ec65020fe7102037fb7b276c9cd121efe2ca40cd Mon Sep 17 00:00:00 2001 From: skydoorkai Date: Mon, 11 Aug 2025 14:24:07 +0800 Subject: [PATCH 2/4] fix tests --- atorch/tests/common_tests/dump_snapshot_test.py | 7 ++++--- atorch/tests/common_tests/trainer/trainer_v2_test.py | 5 ++++- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/atorch/tests/common_tests/dump_snapshot_test.py b/atorch/tests/common_tests/dump_snapshot_test.py index 79adc05..c29c703 100644 --- a/atorch/tests/common_tests/dump_snapshot_test.py +++ b/atorch/tests/common_tests/dump_snapshot_test.py @@ -33,12 +33,9 @@ from atorch.common.util_func import find_free_port # noqa: E402 from atorch.trainer.args import AtorchTrainingArgs # noqa: E402 from atorch.trainer.atorch_trainer_v2 import AtorchTrainerV2 # noqa: E402 -from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 from atorch.utils.import_util import is_coverage_available, is_megatron_lm_available # noqa: E402 from atorch.utils.version import is_megatron_version_bigger_than, torch_version # noqa: E402 -assert is_megatron_lm_available(), f"Can't import megatron, PYTHONPATH={os.environ['PYTHONPATH']}" - if is_megatron_lm_available(): import megatron.legacy.model from megatron.core import mpu @@ -60,6 +57,10 @@ ) from megatron.training.yaml_arguments import core_transformer_config_from_yaml + from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 +else: + pytest.skip("rmegatron not available.", allow_module_level=True) + if is_coverage_available(): import coverage diff --git a/atorch/tests/common_tests/trainer/trainer_v2_test.py b/atorch/tests/common_tests/trainer/trainer_v2_test.py index cbf0546..04fab38 100644 --- a/atorch/tests/common_tests/trainer/trainer_v2_test.py +++ b/atorch/tests/common_tests/trainer/trainer_v2_test.py @@ -22,7 +22,6 @@ from atorch.common.util_func import find_free_port # noqa: E402 from atorch.trainer.args import AtorchTrainingArgs # noqa: E402 from atorch.trainer.atorch_trainer_v2 import AtorchTrainerV2 # noqa: E402 -from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 from atorch.trainer.utils import DistributedType # noqa: E402 from atorch.utils.import_util import is_megatron_lm_available # noqa: E402 from atorch.utils.version import is_megatron_version_bigger_than, torch_version # noqa: E402 @@ -50,6 +49,10 @@ ) from megatron.training.yaml_arguments import core_transformer_config_from_yaml + from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 +else: + pytest.skip("rmegatron not available.", allow_module_level=True) + def model_provider(pre_process=True, post_process=True) -> Union["GPTModel", "megatron.legacy.model.GPTModel"]: """Builds the model. From 96059c5c89bbfb6921c03f274f07d4cb4240f1dd Mon Sep 17 00:00:00 2001 From: skydoorkai Date: Mon, 11 Aug 2025 14:26:37 +0800 Subject: [PATCH 3/4] fix tests --- atorch/tests/common_tests/trainer/trainer_v2_test.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/atorch/tests/common_tests/trainer/trainer_v2_test.py b/atorch/tests/common_tests/trainer/trainer_v2_test.py index 04fab38..9f1638c 100644 --- a/atorch/tests/common_tests/trainer/trainer_v2_test.py +++ b/atorch/tests/common_tests/trainer/trainer_v2_test.py @@ -26,8 +26,6 @@ from atorch.utils.import_util import is_megatron_lm_available # noqa: E402 from atorch.utils.version import is_megatron_version_bigger_than, torch_version # noqa: E402 -assert is_megatron_lm_available(), f"Can't import megatron, PYTHONPATH={os.environ['PYTHONPATH']}" - if is_megatron_lm_available(): import megatron.legacy.model from megatron.core import mpu From 868c984d9ef05fef5a7af77d093f0ee32d871e33 Mon Sep 17 00:00:00 2001 From: skydoorkai Date: Mon, 11 Aug 2025 14:38:05 +0800 Subject: [PATCH 4/4] update ut --- atorch/tests/common_tests/dump_snapshot_test.py | 2 +- .../tests/common_tests/local_sgd_megatron_trainer_test.py | 7 ++++--- .../tests/common_tests/trainer/megatron_dataloader_test.py | 2 ++ atorch/tests/common_tests/trainer/trainer_v2_test.py | 2 +- 4 files changed, 8 insertions(+), 5 deletions(-) diff --git a/atorch/tests/common_tests/dump_snapshot_test.py b/atorch/tests/common_tests/dump_snapshot_test.py index c29c703..6cf56e8 100644 --- a/atorch/tests/common_tests/dump_snapshot_test.py +++ b/atorch/tests/common_tests/dump_snapshot_test.py @@ -59,7 +59,7 @@ from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 else: - pytest.skip("rmegatron not available.", allow_module_level=True) + pytest.skip("megatron not available.", allow_module_level=True) if is_coverage_available(): import coverage diff --git a/atorch/tests/common_tests/local_sgd_megatron_trainer_test.py b/atorch/tests/common_tests/local_sgd_megatron_trainer_test.py index 9e3fca0..3d9af1a 100644 --- a/atorch/tests/common_tests/local_sgd_megatron_trainer_test.py +++ b/atorch/tests/common_tests/local_sgd_megatron_trainer_test.py @@ -28,12 +28,9 @@ from atorch.common.util_func import find_free_port # noqa: E402 from atorch.trainer.args import AtorchTrainingArgs # noqa: E402 from atorch.trainer.atorch_trainer_v2 import AtorchTrainerV2 # noqa: E402 -from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 from atorch.utils.import_util import is_megatron_lm_available # noqa: E402 from atorch.utils.version import get_megatron_version, is_megatron_version_bigger_than, torch_version # noqa: E402 -assert is_megatron_lm_available(), f"Can't import megatron, PYTHONPATH={os.environ['PYTHONPATH']}" - if is_megatron_lm_available(): import megatron.legacy.model from megatron.core import mpu @@ -54,6 +51,10 @@ ) from megatron.training.yaml_arguments import core_transformer_config_from_yaml + from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 +else: + pytest.skip("megatron not available.", allow_module_level=True) + def model_provider(pre_process=True, post_process=True) -> Union["GPTModel", "megatron.legacy.model.GPTModel"]: """Builds the model. diff --git a/atorch/tests/common_tests/trainer/megatron_dataloader_test.py b/atorch/tests/common_tests/trainer/megatron_dataloader_test.py index 9e94b29..12f55e7 100644 --- a/atorch/tests/common_tests/trainer/megatron_dataloader_test.py +++ b/atorch/tests/common_tests/trainer/megatron_dataloader_test.py @@ -28,6 +28,8 @@ MegatronDataloaderWrapper, skip_first_batches_for_megatron_dataloader, ) +else: + pytest.skip("megatron not available.", allow_module_level=True) class DummyDataset(Dataset): diff --git a/atorch/tests/common_tests/trainer/trainer_v2_test.py b/atorch/tests/common_tests/trainer/trainer_v2_test.py index 9f1638c..055d216 100644 --- a/atorch/tests/common_tests/trainer/trainer_v2_test.py +++ b/atorch/tests/common_tests/trainer/trainer_v2_test.py @@ -49,7 +49,7 @@ from atorch.trainer.megatron import MegatronTrainStep # noqa: E402 else: - pytest.skip("rmegatron not available.", allow_module_level=True) + pytest.skip("megatron not available.", allow_module_level=True) def model_provider(pre_process=True, post_process=True) -> Union["GPTModel", "megatron.legacy.model.GPTModel"]: