diff --git a/megatron/core/ssm/gated_delta_net.py b/megatron/core/ssm/gated_delta_net.py index 391bd3c7977..f56754012aa 100644 --- a/megatron/core/ssm/gated_delta_net.py +++ b/megatron/core/ssm/gated_delta_net.py @@ -148,29 +148,62 @@ def __init__( # Input projection (hidden_states -> q, k, v, gate, beta, alpha) # TODO: for now, output gate is forced for GDN. # We may remove this restriction in the future. - self.in_proj_dim = self.qk_dim * 2 + self.v_dim * 2 + self.num_value_heads * 2 + self.conv_dim = self.qk_dim * 2 + self.v_dim + self.gba_dim = self.v_dim + self.num_value_heads * 2 + self.gba_dim_local_tp = self.gba_dim // self.tp_size + self.in_proj_dim = self.conv_dim + self.gba_dim + self.split_in_proj = os.environ.get("MCORE_GDN_SPLIT_IN_PROJ", "0") == "1" if self.config.fp8: fp8_align_size = get_fp8_align_size(self.config.fp8_recipe) - assert self.in_proj_dim % fp8_align_size == 0, ( - "For FP8, the innermost dimension of the GDN layer " - "input projection output tensor must be a multiple of 16." + fp8_dims = [self.conv_dim, self.gba_dim] if self.split_in_proj else [self.in_proj_dim] + for fp8_dim in fp8_dims: + assert fp8_dim % fp8_align_size == 0, ( + "For FP8, the innermost dimension of each GDN layer " + "input projection output tensor must be a multiple of 16." + ) + if self.split_in_proj: + self.qkv_proj = build_module( + submodules.in_proj, + self.hidden_size, + self.conv_dim, + config=self.config, + init_method=self.config.init_method, + gather_output=False, + bias=bias, + skip_bias_add=False, + is_expert=False, + tp_comm_buffer_name="fc1_qkv", + tp_group=self.pg_collection.tp, + ) + self.gba_proj = build_module( + submodules.in_proj, + self.hidden_size, + self.gba_dim, + config=self.config, + init_method=self.config.init_method, + gather_output=False, + bias=bias, + skip_bias_add=False, + is_expert=False, + tp_comm_buffer_name="fc1_gba", + tp_group=self.pg_collection.tp, + ) + else: + self.in_proj = build_module( + submodules.in_proj, + self.hidden_size, + self.in_proj_dim, + config=self.config, + init_method=self.config.init_method, + gather_output=False, + bias=bias, + skip_bias_add=False, + is_expert=False, + tp_comm_buffer_name="fc1", + tp_group=self.pg_collection.tp, ) - self.in_proj = build_module( - submodules.in_proj, - self.hidden_size, - self.in_proj_dim, - config=self.config, - init_method=self.config.init_method, - gather_output=False, - bias=bias, - skip_bias_add=False, - is_expert=False, - tp_comm_buffer_name="fc1", - tp_group=self.pg_collection.tp, - ) # Conv1d for QKV - self.conv_dim = self.qk_dim * 2 + self.v_dim self.conv_dim_local_tp = self.conv_dim // self.tp_size # weight shape: [conv_dim, 1, d_conv] @@ -339,63 +372,125 @@ def forward( cu_seqlens_q = None cu_seqlens_kv = None - # Input projection - nvtx_range_push(suffix="in_proj") - qkvzba, _ = self.in_proj(hidden_states) - nvtx_range_pop(suffix="in_proj") - - # CP All to All: CP to HP - if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': - unpacked_qkvzba = _unpack_sequence(qkvzba, cu_seqlens_q // self.cp_size, dim=0) - outputs = [] - for qkvzba_i in unpacked_qkvzba: - qkvzba_i = tensor_a2a_cp2hp( - qkvzba_i, + qkv_channels_split_sections = [ + self.qk_dim_local_tp, + self.qk_dim_local_tp, + self.v_dim_local_tp, + ] + gba_split_sections = [ + self.v_dim_local_tp, + self.num_value_heads // self.tp_size, + self.num_value_heads // self.tp_size, + ] + if self.split_in_proj: + # Input projection + nvtx_range_push(suffix="qkv_proj") + qkv, _ = self.qkv_proj(hidden_states) + nvtx_range_pop(suffix="qkv_proj") + nvtx_range_push(suffix="gba_proj") + gba, _ = self.gba_proj(hidden_states) + nvtx_range_pop(suffix="gba_proj") + + # CP All to All: CP to HP + if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': + unpacked_qkv = _unpack_sequence(qkv, cu_seqlens_q // self.cp_size, dim=0) + unpacked_gba = _unpack_sequence(gba, cu_seqlens_q // self.cp_size, dim=0) + qkv_outputs = [] + gba_outputs = [] + for qkv_i, gba_i in zip(unpacked_qkv, unpacked_gba): + qkv_outputs.append( + tensor_a2a_cp2hp( + qkv_i, + seq_dim=0, + head_dim=-1, + cp_group=self.pg_collection.cp, + split_sections=qkv_channels_split_sections, + ) + ) + gba_outputs.append( + tensor_a2a_cp2hp( + gba_i, + seq_dim=0, + head_dim=-1, + cp_group=self.pg_collection.cp, + split_sections=gba_split_sections, + ) + ) + qkv = torch.cat(qkv_outputs, dim=0) + gba = torch.cat(gba_outputs, dim=0) + else: + qkv = tensor_a2a_cp2hp( + qkv, + seq_dim=0, + head_dim=-1, + cp_group=self.pg_collection.cp, + split_sections=qkv_channels_split_sections, + ) + gba = tensor_a2a_cp2hp( + gba, seq_dim=0, head_dim=-1, cp_group=self.pg_collection.cp, - split_sections=[ - self.qk_dim_local_tp, - self.qk_dim_local_tp, - self.v_dim_local_tp, - self.v_dim_local_tp, - self.num_value_heads // self.tp_size, - self.num_value_heads // self.tp_size, - ], + split_sections=gba_split_sections, ) - outputs.append(qkvzba_i) - qkvzba = torch.cat(outputs, dim=0) + + # Transpose: s b x --> b s x + # From sbhd to bshd format + qkv = qkv.transpose(0, 1) + gba = gba.transpose(0, 1) + gate, beta, alpha = torch.split( + gba, + [ + self.v_dim_local_tp // self.cp_size, + self.num_value_heads // self.tp_size // self.cp_size, + self.num_value_heads // self.tp_size // self.cp_size, + ], + dim=-1, + ) else: - qkvzba = tensor_a2a_cp2hp( + # Input projection + nvtx_range_push(suffix="in_proj") + qkvzba, _ = self.in_proj(hidden_states) + nvtx_range_pop(suffix="in_proj") + + # CP All to All: CP to HP + if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': + unpacked_qkvzba = _unpack_sequence(qkvzba, cu_seqlens_q // self.cp_size, dim=0) + outputs = [] + for qkvzba_i in unpacked_qkvzba: + qkvzba_i = tensor_a2a_cp2hp( + qkvzba_i, + seq_dim=0, + head_dim=-1, + cp_group=self.pg_collection.cp, + split_sections=qkv_channels_split_sections + gba_split_sections, + ) + outputs.append(qkvzba_i) + qkvzba = torch.cat(outputs, dim=0) + else: + qkvzba = tensor_a2a_cp2hp( + qkvzba, + seq_dim=0, + head_dim=-1, + cp_group=self.pg_collection.cp, + split_sections=qkv_channels_split_sections + gba_split_sections, + ) + + # Transpose: s b x --> b s x + # From sbhd to bshd format + qkvzba = qkvzba.transpose(0, 1) + + # Split, reorder, and reshape the tensor into q, k, v, gate, beta, alpha + qkv, gate, beta, alpha = torch.split( qkvzba, - seq_dim=0, - head_dim=-1, - cp_group=self.pg_collection.cp, - split_sections=[ - self.qk_dim_local_tp, - self.qk_dim_local_tp, - self.v_dim_local_tp, - self.v_dim_local_tp, - self.num_value_heads // self.tp_size, - self.num_value_heads // self.tp_size, + [ + self.conv_dim_local_tp // self.cp_size, + self.v_dim_local_tp // self.cp_size, + self.num_value_heads // self.tp_size // self.cp_size, + self.num_value_heads // self.tp_size // self.cp_size, ], + dim=-1, ) - - # Transpose: s b x --> b s x - # From sbhd to bshd format - qkvzba = qkvzba.transpose(0, 1) - - # Split, reorder, and reshape the tensor into q, k, v, gate, beta, alpha - qkv, gate, beta, alpha = torch.split( - qkvzba, - [ - (self.qk_dim_local_tp * 2 + self.v_dim_local_tp) // self.cp_size, - self.v_dim_local_tp // self.cp_size, - self.num_value_heads // self.tp_size // self.cp_size, - self.num_value_heads // self.tp_size // self.cp_size, - ], - dim=-1, - ) gate = gate.reshape(batch, seq_len, -1, self.value_head_dim) beta = beta.reshape(batch, seq_len, -1) alpha = alpha.reshape(batch, seq_len, -1) @@ -403,11 +498,6 @@ def forward( # Convolution on qkv nvtx_range_push(suffix="conv1d") seq_len = qkv.shape[1] - qkv_channels_split_sections = [ - self.qk_dim_local_tp, - self.qk_dim_local_tp, - self.v_dim_local_tp, - ] conv1d_weight = get_parameter_local_cp( self.conv1d.weight, dim=0, @@ -637,25 +727,54 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None, tp_gr # At this point the TP sharding is correctly defined for each tensor, but some of the # tensors must be additionally split into separate parts - in_proj_dim_local_tp = self.in_proj_dim // self.tp_size - assert sharded_state_dict[f"{prefix}in_proj.weight"].data.size(0) == in_proj_dim_local_tp, ( - in_proj_dim_local_tp, - sharded_state_dict[f"{prefix}in_proj.weight"], - ) + if self.split_in_proj: + assert ( + sharded_state_dict[f"{prefix}qkv_proj.weight"].data.size(0) + == self.conv_dim_local_tp + ), (self.conv_dim_local_tp, sharded_state_dict[f"{prefix}qkv_proj.weight"]) + assert ( + sharded_state_dict[f"{prefix}gba_proj.weight"].data.size(0) + == self.gba_dim_local_tp + ), (self.gba_dim_local_tp, sharded_state_dict[f"{prefix}gba_proj.weight"]) - sharded_state_dict[f"{prefix}in_proj.weight"] = _split_tensor_factory( - sharded_state_dict[f"{prefix}in_proj.weight"], - [ - self.qk_dim_local_tp, - self.qk_dim_local_tp, - self.v_dim_local_tp, - self.v_dim_local_tp, - self.num_value_heads // self.tp_size, - self.num_value_heads // self.tp_size, - ], - ["query", "key", "value", "z", "beta", "alpha"], - 0, - ) + sharded_state_dict[f"{prefix}qkv_proj.weight"] = _split_tensor_factory( + sharded_state_dict[f"{prefix}qkv_proj.weight"], + [self.qk_dim_local_tp, self.qk_dim_local_tp, self.v_dim_local_tp], + ["query", "key", "value"], + 0, + ) + sharded_state_dict[f"{prefix}gba_proj.weight"] = _split_tensor_factory( + sharded_state_dict[f"{prefix}gba_proj.weight"], + [ + self.v_dim_local_tp, + self.num_value_heads // self.tp_size, + self.num_value_heads // self.tp_size, + ], + ["z", "beta", "alpha"], + 0, + ) + else: + in_proj_dim_local_tp = self.in_proj_dim // self.tp_size + assert sharded_state_dict[f"{prefix}in_proj.weight"].data.size( + 0 + ) == in_proj_dim_local_tp, ( + in_proj_dim_local_tp, + sharded_state_dict[f"{prefix}in_proj.weight"], + ) + + sharded_state_dict[f"{prefix}in_proj.weight"] = _split_tensor_factory( + sharded_state_dict[f"{prefix}in_proj.weight"], + [ + self.qk_dim_local_tp, + self.qk_dim_local_tp, + self.v_dim_local_tp, + self.v_dim_local_tp, + self.num_value_heads // self.tp_size, + self.num_value_heads // self.tp_size, + ], + ["query", "key", "value", "z", "beta", "alpha"], + 0, + ) conv_layer_name_list = ["conv1d.weight"] assert ( @@ -678,13 +797,21 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None, tp_gr def backward_dw(self): """Execute weight gradient computation for all linear layers.""" - self._backward_in_proj() + if self.split_in_proj: + self._backward_split_in_proj() + else: + self._backward_in_proj() self._backward_out_proj() def _backward_in_proj(self): """Computes weight gradients of input projection layer.""" self.in_proj.backward_dw() + def _backward_split_in_proj(self): + """Computes weight gradients of split input projection layers.""" + self.qkv_proj.backward_dw() + self.gba_proj.backward_dw() + def _backward_out_proj(self): """Computes weight gradients of output projection layer.""" self.out_proj.backward_dw() diff --git a/tests/unit_tests/ssm/bench_gdn_cuda_opt.py b/tests/unit_tests/ssm/bench_gdn_cuda_opt.py index 9668f18a4d1..36040b936da 100644 --- a/tests/unit_tests/ssm/bench_gdn_cuda_opt.py +++ b/tests/unit_tests/ssm/bench_gdn_cuda_opt.py @@ -26,6 +26,7 @@ from tests.unit_tests.test_utilities import Utils FLAGS = ( + "MCORE_GDN_SPLIT_IN_PROJ", "MCORE_GDN_USE_OPT_WRAPPER", "MCORE_GDN_OPT_BACKEND", "MCORE_GDN_OPT_WARN_FALLBACK", @@ -203,9 +204,14 @@ def scenario_label(index, name): return f"gdn_only/{index:02d}_{safe_name}" -def make_model(dtype): +def make_model(dtype, split_in_proj=False): from megatron.core.ssm.gated_delta_net import GatedDeltaNet + if split_in_proj: + os.environ["MCORE_GDN_SPLIT_IN_PROJ"] = "1" + else: + os.environ.pop("MCORE_GDN_SPLIT_IN_PROJ", None) + Utils.initialize_model_parallel( tensor_model_parallel_size=1, pipeline_model_parallel_size=1, context_parallel_size=1 ) @@ -397,6 +403,7 @@ def parse_args(): parser.add_argument("--rtol", type=float, default=5e-3) parser.add_argument("--fail-on-accuracy", action="store_true") parser.add_argument("--no-nvtx", dest="use_nvtx", action="store_false", default=True) + parser.add_argument("--split-in-proj", action="store_true") return parser.parse_args() @@ -416,10 +423,10 @@ def main(): set_env({}) print( f"DEVICE {torch.cuda.get_device_name(0)} SHAPE B=2 T=8192 H=64 D=128 " - f"dtype={args.dtype} loss={args.loss}" + f"dtype={args.dtype} loss={args.loss} split_in_proj={args.split_in_proj}" ) try: - model = make_model(dtype).eval() + model = make_model(dtype, split_in_proj=args.split_in_proj).eval() x = torch.randn(8192, 2, 128, device="cuda", dtype=dtype) accuracy_rows = check_accuracy( model, x, scenario_items, args.loss, args.atol, args.rtol, args.use_nvtx diff --git a/tests/unit_tests/ssm/test_bench_gdn_cuda_opt_scenarios.py b/tests/unit_tests/ssm/test_bench_gdn_cuda_opt_scenarios.py index 0f57245ecdf..583889a7d89 100644 --- a/tests/unit_tests/ssm/test_bench_gdn_cuda_opt_scenarios.py +++ b/tests/unit_tests/ssm/test_bench_gdn_cuda_opt_scenarios.py @@ -54,3 +54,12 @@ def test_benchmark_does_not_expose_dhu_dqkwg_wrapper_path(): assert forbidden.isdisjoint(flags) for key, (_label, env) in scenarios.items(): assert forbidden.isdisjoint(env), key + + +def test_split_in_proj_is_construction_flag_not_scenario_flag(): + scenarios = _literal_assignment("SCENARIOS") + flags = _literal_assignment("FLAGS") + + assert "MCORE_GDN_SPLIT_IN_PROJ" in flags + for key, (_label, env) in scenarios.items(): + assert "MCORE_GDN_SPLIT_IN_PROJ" not in env, key diff --git a/tests/unit_tests/ssm/test_gated_delta_net.py b/tests/unit_tests/ssm/test_gated_delta_net.py index 3eb02442fe9..2153e70b806 100644 --- a/tests/unit_tests/ssm/test_gated_delta_net.py +++ b/tests/unit_tests/ssm/test_gated_delta_net.py @@ -348,6 +348,76 @@ def test_cp1_still_validates_total(self, mock_gdn): mock_gdn._resolve_cu_seqlens(None, actual, 1008, "cu_seqlens_q") +@pytest.mark.skipif(not HAVE_FLA, reason="FLA is not installed.") +@pytest.mark.internal +def test_gated_delta_net_split_in_proj_builds_expected_projections(monkeypatch): + monkeypatch.setenv("MCORE_GDN_SPLIT_IN_PROJ", "1") + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=1, + ) + try: + model_parallel_cuda_manual_seed(123) + pg_collection = ProcessGroupCollection( + tp=parallel_state.get_tensor_model_parallel_group(), + cp=parallel_state.get_context_parallel_group(), + ) + transformer_config = TransformerConfig( + hidden_size=256, + linear_conv_kernel_dim=2, + linear_key_head_dim=64, + linear_value_head_dim=64, + linear_num_key_heads=4, + linear_num_value_heads=8, + num_layers=1, + normalization="RMSNorm", + use_cpu_initialization=True, + layernorm_zero_centered_gamma=True, + num_attention_heads=8, + activation_func=F.silu, + bf16=True, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=[1], + transformer_impl="transformer_engine", + ) + gdn_submodules = get_experimental_attention_variant_module_spec( + config=transformer_config + ).submodules + gdn = ( + GatedDeltaNet( + transformer_config, + submodules=gdn_submodules, + layer_number=1, + bias=False, + conv_bias=False, + conv_init=1.0, + use_qk_l2norm=True, + A_init_range=(1, 16), + pg_collection=pg_collection, + ) + .cuda() + .bfloat16() + ) + + assert gdn.split_in_proj + assert not hasattr(gdn, "in_proj") + assert gdn.qkv_proj.weight.shape[0] == gdn.conv_dim_local_tp + assert gdn.gba_proj.weight.shape[0] == gdn.gba_dim_local_tp + + hidden_states = torch.randn( + 16, 2, gdn.config.hidden_size, device=torch.cuda.current_device(), dtype=torch.bfloat16 + ) + output, _ = gdn(hidden_states, attention_mask=None) + output.float().square().mean().backward() + + assert output.shape == hidden_states.shape + assert gdn.qkv_proj.weight.grad is not None + assert gdn.gba_proj.weight.grad is not None + finally: + Utils.destroy_model_parallel() + + @pytest.mark.parametrize("sequence_packing", [False, True]) @pytest.mark.parametrize( ("tp", "sp", "cp"),