# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 from abc import ABC, abstractmethod import torch import torch.distributed as dist from megatron.core import mpu, tensor_parallel from megatron.core.inference.contexts import BaseInferenceContext from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.transformer.module import MegatronModule from transformers import AutoConfig class HuggingfaceAttention(MegatronModule, ABC): """Attention layer abstract class. This layer only contains common modules required for the "self attn" and "cross attn" specializations. """ def __init__( self, args, config, layer_number: int, cp_comm_type: str = "p2p", pg_collection=None, ): super().__init__(config=config) self.args = args self.config = config # Note that megatron layer_number starts at 1 self.layer_number = layer_number self.hf_layer_idx = layer_number - 1 self.hf_config = AutoConfig.from_pretrained(args.hf_checkpoint, trust_remote_code=True) # hardcode to fa2 at the moment. self.hf_config._attn_implementation = "flash_attention_2" def forward( self, hidden_states: torch.Tensor, attention_mask: torch.Tensor, key_value_states: torch.Tensor | None = None, inference_context: BaseInferenceContext | None = None, rotary_pos_emb: torch.Tensor | tuple[torch.Tensor, torch.Tensor] | None = None, rotary_pos_cos: torch.Tensor | None = None, rotary_pos_sin: torch.Tensor | None = None, rotary_pos_cos_sin: torch.Tensor | None = None, attention_bias: torch.Tensor | None = None, packed_seq_params: PackedSeqParams | None = None, sequence_len_offset: int | None = None, *, inference_params: BaseInferenceContext | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: assert packed_seq_params is not None cu_seqlens = packed_seq_params.cu_seqlens_q if self.args.sequence_parallel: hidden_states = tensor_parallel.gather_from_sequence_parallel_region( hidden_states, group=mpu.get_tensor_model_parallel_group() ) if mpu.get_context_parallel_world_size() > 1: cp_size = mpu.get_context_parallel_world_size() hidden_states_list = dist.nn.all_gather( hidden_states, group=mpu.get_context_parallel_group(), ) # TODO: preprocess this for each batch to prevent tolist in the training step whole_hidden_states_list = [] local_cu_seqlens = cu_seqlens // cp_size for i in range(len(cu_seqlens) - 1): seqlen = cu_seqlens[i + 1] - cu_seqlens[i] chunk_size = seqlen // 2 // cp_size whole_hidden_states_list.extend( [ hidden_states_list[cp_rank][local_cu_seqlens[i] : local_cu_seqlens[i] + chunk_size] for cp_rank in range(cp_size) ] + [ hidden_states_list[cp_rank][local_cu_seqlens[i] + chunk_size : local_cu_seqlens[i + 1]] for cp_rank in range(cp_size) ][::-1], ) hidden_states = torch.cat(whole_hidden_states_list, dim=0) hidden_states = hidden_states.permute(1, 0, 2) # [bsz, seq_len, hidden_dim] output = self.hf_forward(hidden_states, packed_seq_params) bias = None output = output.permute(1, 0, 2) # [seq_len, bsz, hidden_dim] if mpu.get_context_parallel_world_size() > 1: cp_rank = mpu.get_context_parallel_rank() output_list = [] for i in range(len(cu_seqlens) - 1): seqlen = cu_seqlens[i + 1] - cu_seqlens[i] chunk_size = seqlen // 2 // cp_size seq = output[cu_seqlens[i] : cu_seqlens[i + 1]] chunks = torch.chunk(seq, 2 * cp_size, dim=0) output_list.append(chunks[cp_rank]) output_list.append(chunks[2 * cp_size - 1 - cp_rank]) output = torch.cat(output_list, dim=0) if self.args.sequence_parallel: output = tensor_parallel.scatter_to_sequence_parallel_region( output, group=mpu.get_tensor_model_parallel_group() ) return output, bias @abstractmethod def hf_forward(self, hidden_states, packed_seq_params): """Huggingface forward function"""