Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Empty file added d2/planner/ilp_planner.py
Empty file.
199 changes: 197 additions & 2 deletions d2/runtime/megatron/base_transformer_layer.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import Any, Optional
from typing import Any, Iterable, Optional
import time
import torch
from torch import Tensor
Expand All @@ -13,6 +13,7 @@
TransformerLayer as MegatronTransformerLayer,
TransformerLayerSubmodules,
)
from megatron.core import parallel_state

from d2.runtime.megatron.packed_seq_params import PingPangSingleStepPackedSeqParams

Expand All @@ -38,6 +39,201 @@ def __init__(
from megatron.core.extensions.transformer_engine import TEDotProductAttention
self.self_attention: SelfAttention
assert isinstance(self.self_attention.core_attention, TEDotProductAttention)
self.pre_attn_cuda_graph = None
self.post_attn_cuda_graph = None

def init_pre_attn_cuda_graph(self, prev_layer: "Optional[TransformerLayer]", seq_len: int, device: torch.device, dtype: torch.dtype):
# TODO: May need to also check `self.config.sequence_parallel_size`. If not, just assume SP == TP
# use_sp = self.config.sequence_parallel
tp = parallel_state.get_tensor_model_parallel_world_size()
if prev_layer is None:
static_input = torch.zeros((seq_len, 1, self.config.hidden_size), device=device, dtype=dtype, requires_grad=True)
self.pre_attn_cuda_graph = torch.cuda.make_graphed_callables(self.__pre_attn_cuda_graph, (static_input,))
else:
hidden_size_tp = self.config.hidden_size // tp
# prev_layer.self_attention.query_projection_size // tp
static_core_attn_out = torch.zeros((seq_len * tp, 1, hidden_size_tp), device=device, dtype=dtype, requires_grad=True)
# static_residual = torch.zeros((seq_len // tp, 1, self.config.hidden_size), device=device, dtype=dtype, requires_grad=True)
static_residual = torch.zeros((seq_len, 1, self.config.hidden_size), device=device, dtype=dtype, requires_grad=True)
def post_then_pre_core_attn_cuda_graph(core_attn_out: Tensor, residual: Tensor):
hidden_states = prev_layer._post_attn_cuda_graph(core_attn_out, residual)
return self.__pre_attn_cuda_graph(hidden_states)
self.pre_attn_cuda_graph = torch.cuda.make_graphed_callables(post_then_pre_core_attn_cuda_graph, (static_core_attn_out, static_residual))

def init_post_attn_cuda_graph(self, seq_len: int, device: torch.device, dtype: torch.dtype):
tp = parallel_state.get_tensor_model_parallel_world_size()
hidden_size_tp = self.config.hidden_size // tp
# prev_layer.self_attention.query_projection_size // tp
static_core_attn_out = torch.zeros((seq_len * tp, 1, hidden_size_tp), device=device, dtype=dtype, requires_grad=True)
static_residual = torch.zeros((seq_len, 1, self.config.hidden_size), device=device, dtype=dtype, requires_grad=True)
self.post_attn_cuda_graph = torch.cuda.make_graphed_callables(self._post_attn_cuda_graph, (static_core_attn_out, static_residual))

def _forward_pre_attn_cuda_graph(
self,
args: Iterable[torch.Tensor],
rotary_pos_emb: Optional[Tensor] = None,
rotary_pos_cos: Optional[Tensor] = None,
rotary_pos_sin: Optional[Tensor] = None,
packed_seq_params: Optional[PingPangSingleStepPackedSeqParams] = None,
sequence_len_offset: Optional[Tensor] = None
):
query, key, value, residual = self.pre_attn_cuda_graph(*args)

log_memory_usage(f"(L{self.layer_number}) _forward_pre_core_attn:(init, before input layernorm)")

assert rotary_pos_cos is None and rotary_pos_sin is None

# For self attention we just duplicate the rotary_pos_emb if it isn't already
if rotary_pos_emb is not None and not isinstance(rotary_pos_emb, tuple):
rotary_pos_emb = (rotary_pos_emb,) * 2

#### Some code in core_attention. This is because we don't want the pos embedding
# being handled in the attention layout (the pos id will be hard to handle)
inference_context = None

log_memory_usage(f"(L{self.layer_number}) _forward_pre_core_attn:(after qkv, before adjust_key_value_for_inference)")

query, key, value, rotary_pos_emb, attn_mask_type = self.self_attention._adjust_key_value_for_inference(
inference_context,
query,
key,
value,
rotary_pos_emb,
rotary_pos_cos,
rotary_pos_sin,
sequence_len_offset,
)
if packed_seq_params is not None:
query = query.squeeze(1)
key = key.squeeze(1)
value = value.squeeze(1)

log_memory_usage(f"(L{self.layer_number}) _forward_pre_core_attn:(after adjust_key_value_for_inference, before rope)")

# ================================================
# relative positional embedding (rotary embedding)
# ================================================
if rotary_pos_emb is not None and not self.config.flash_decode:
q_pos_emb, k_pos_emb = rotary_pos_emb

if packed_seq_params is not None:
if packed_seq_params.cu_seqlens_q_padded is not None:
cu_seqlens_q = packed_seq_params.cu_seqlens_q_padded
else:
cu_seqlens_q = packed_seq_params.cu_seqlens_q
if packed_seq_params.cu_seqlens_kv_padded is not None:
cu_seqlens_kv = packed_seq_params.cu_seqlens_kv_padded
else:
cu_seqlens_kv = packed_seq_params.cu_seqlens_kv
else:
cu_seqlens_q = cu_seqlens_kv = None

if q_pos_emb is not None:
# TODO VIJAY: simplify
query = apply_rotary_pos_emb(
query, q_pos_emb, config=self.config, cu_seqlens=cu_seqlens_q
)
if k_pos_emb is not None:
key = apply_rotary_pos_emb(
key, k_pos_emb, config=self.config, cu_seqlens=cu_seqlens_kv
)

# TODO, can apply positional embedding to value_layer so it has
# absolute positional embedding.
# otherwise, only relative positional embedding takes effect
# value_layer = apply_rotary_pos_emb(value_layer, k_pos_emb)
log_memory_usage(f"(L{self.layer_number}) _forward_pre_core_attn:(after rope, before return)")
return query, key, value, residual, attn_mask_type

def _forward_post_attn_cuda_graph(
self, core_attn_out: Tensor, residual: Tensor,
context: Optional[Tensor] = None, context_mask: Optional[Tensor] = None,
):
assert context is None and context_mask is None, "not supported in cudagraph"
return self.post_attn_cuda_graph(core_attn_out, residual)


def __pre_attn_cuda_graph(self, hidden_states: Tensor):
if self.recompute_input_layernorm:
self.input_layernorm_checkpoint = tensor_parallel.CheckpointWithoutOutput()
input_layernorm_output = self.input_layernorm_checkpoint.checkpoint(
self.input_layernorm, hidden_states
)
else:
input_layernorm_output = self.input_layernorm(hidden_states)

log_memory_usage(f"(L{self.layer_number}) _forward_pre_core_attn:(after input layernorm, before qkv)")

# q, k, v
query, key, value = self.self_attention.get_query_key_value_tensors(input_layernorm_output, None)
return query, key, value, hidden_states

def _post_attn_cuda_graph(self, core_attn_out: Tensor, residual: Tensor):
attention_output_with_bias = self.self_attention.linear_proj(core_attn_out)

log_memory_usage(f"(L{self.layer_number}) _forward_post_core_attn:(before layernorm)")
if self.recompute_input_layernorm:
# discard the output of the input layernorm and register the recompute
# as a gradient hook of attention_output_with_bias[0]
self.input_layernorm_checkpoint.discard_output_and_register_recompute(
attention_output_with_bias[0]
)

log_memory_usage(f"(L{self.layer_number}) _forward_post_core_attn:(after layernorm)")

# TODO: could we move `bias_dropout_add_exec_handler` itself
# inside the module provided in the `bias_dropout_add_spec` module?
with self.bias_dropout_add_exec_handler():
hidden_states = self.self_attn_bda(self.training, self.config.bias_dropout_fusion)(
attention_output_with_bias, residual, self.hidden_dropout
)

# Residual connection.
residual = hidden_states

# # Optional Layer norm after self-attention
# pre_cross_attn_layernorm_output = self.pre_cross_attn_layernorm(hidden_states)

# # Cross attention.
# log_memory_usage(f"(L{self.layer_number}) _forward_post_core_attn:(cross attention)")
# attention_output_with_bias = self.cross_attention(
# pre_cross_attn_layernorm_output,
# attention_mask=context_mask,
# key_value_states=context,
# inference_context=inference_context,
# )
# log_memory_usage(f"(L{self.layer_number}) _forward_post_core_attn:(after cross attention)")

# if isinstance(attention_output_with_bias, dict) and "context" in attention_output_with_bias:
# context = attention_output_with_bias["context"]

attention_output_with_bias = self.pre_cross_attn_layernorm(hidden_states)

# TODO: could we move `bias_dropout_add_exec_handler` itself
# inside the module provided in the `bias_dropout_add_spec` module?
with self.bias_dropout_add_exec_handler():
hidden_states = self.cross_attn_bda(self.training, self.config.bias_dropout_fusion)(
attention_output_with_bias, residual, self.hidden_dropout
)
log_memory_usage(f"(L{self.layer_number}) _forward_post_core_attn:(after cross attn bda)")

# Residual connection.
residual = hidden_states

# Optional Layer norm post the cross-attention.
if self.recompute_pre_mlp_layernorm:
self.pre_mlp_norm_checkpoint = tensor_parallel.CheckpointWithoutOutput()
pre_mlp_layernorm_output = self.pre_mlp_norm_checkpoint.checkpoint(
self.pre_mlp_layernorm, hidden_states
)
else:
pre_mlp_layernorm_output = self.pre_mlp_layernorm(hidden_states)

log_memory_usage(f"(L{self.layer_number}) _forward_post_core_attn:(after pre mlp layernorm)")

mlp_output = self._forward_mlp(pre_mlp_layernorm_output, residual)
log_memory_usage(f"(L{self.layer_number}) _forward_post_core_attn:(after mlp)")
return mlp_output

def _forward_pre_core_attn(
self,
Expand Down Expand Up @@ -75,7 +271,6 @@ def _forward_pre_core_attn(
# q, k, v
log_memory_usage(f"(L{self.layer_number}) _forward_pre_core_attn:(before qkv)")
query, key, value = self.self_attention.get_query_key_value_tensors(input_layernorm_output, None)
# print(f"🟡 query: {query.shape} (query.dtype: {query.dtype}), key: {key.shape} (key.dtype: {key.dtype}), value: {value.shape} (value.dtype: {value.dtype})")

#### Some code in core_attention. This is because we don't want the pos embedding
# being handled in the attention layout (the pos id will be hard to handle)
Expand Down
2 changes: 1 addition & 1 deletion d2/runtime/megatron/forward_backward_func.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,7 +90,7 @@ def send_forward_backward_recv_forward_backward(output_tensors, input_tensor_gra
forward_backward_pipelining_without_interleaving_first_run = True


import wlbllm.registry
# import wlbllm.registry
def wlb_swap_next_forward_metadata():
# Call this function before entering forward step.
swap_metadata_fn = wlbllm.registry.get("swap_metadata_fn")
Expand Down
119 changes: 100 additions & 19 deletions d2/runtime/megatron/ping_pong/tick_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,16 +38,12 @@ def forward_pre_core_attn(layer: TransformerLayer, args: Dict[str, Any]):
)
log_memory_usage(f"(L{layer.layer_number}) forward_pre_core_attn:(after pre core attn)")

log_memory_usage(f"(L{layer.layer_number}) forward_pre_core_attn:(before pre mlp to attn)")
signal = layer._pre_mlp_to_attn(query, key, value, args["packed_seq_params"])
log_memory_usage(f"(L{layer.layer_number}) forward_pre_core_attn:(after pre mlp to attn)")

args["query"] = query
args["key"] = key
args["value"] = value
args["residual"] = residual
args["attn_mask_type"] = attn_mask_type
args["signal"] = signal
forward_pre_core_attn_comm(layer, args)


# Record timing if enabled
Expand All @@ -59,6 +55,16 @@ def forward_pre_core_attn(layer: TransformerLayer, args: Dict[str, Any]):
return args


def forward_pre_core_attn_comm(layer: TransformerLayer, args: Dict[str, Any]):
log_memory_usage(f"(L{layer.layer_number}) forward_pre_core_attn:(before pre mlp to attn)")
signal = layer._pre_mlp_to_attn(
args["query"], args["key"], args["value"], args["packed_seq_params"],
)
log_memory_usage(f"(L{layer.layer_number}) forward_pre_core_attn:(after pre mlp to attn)")
args["signal"] = signal
return args


def layout_mlp_to_attn(layer: TransformerLayer, args: Dict[str, Any]):
log_memory_usage(f"(L{layer.layer_number}) layout_mlp_to_attn:(start)")
bwd_resend_qkv = args["packed_seq_params"].bwd_packed_seq_params is not None
Expand Down Expand Up @@ -126,17 +132,7 @@ def layout_attn_to_mlp(layer: TransformerLayer, args: Dict[str, Any]):
return args


def forward_post_core_attn(layer: TransformerLayer, args: Dict[str, Any]):
log_memory_usage(f"(L{layer.layer_number}) forward_post_core_attn:(start)")

# Setup timing if enabled
start_event = None
end_event = None
if os.getenv("UNIFIED_RECORD_TICK_TIMES", "0") == "1" and _current_tick_sample_id is not None:
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
start_event.record()

def forward_post_core_attn_comm(layer: TransformerLayer, args: Dict[str, Any]):
signal = args.pop("signal")
packed_seq_params: PingPangSingleStepPackedSeqParams = args["packed_seq_params"]
bwd_resend_qkv = packed_seq_params.bwd_packed_seq_params is not None
Expand All @@ -151,10 +147,25 @@ def forward_post_core_attn(layer: TransformerLayer, args: Dict[str, Any]):
args.pop("query"), args.pop("key"), args.pop("value")
else:
core_attn_out = layer._post_attn_to_mlp(signal, args["packed_seq_params"])
residual = args.pop("residual")
args["core_attn_out"] = core_attn_out
return args


def forward_post_core_attn(layer: TransformerLayer, args: Dict[str, Any]):
log_memory_usage(f"(L{layer.layer_number}) forward_post_core_attn:(start)")

# Setup timing if enabled
start_event = None
end_event = None
if os.getenv("UNIFIED_RECORD_TICK_TIMES", "0") == "1" and _current_tick_sample_id is not None:
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
start_event.record()

forward_post_core_attn_comm(layer, args)
mlp_output, context = layer._forward_post_core_attn(
core_attn_out,
residual,
args.pop("core_attn_out"),
args.pop("residual"),
args["context"],
args["context_mask"],
)
Expand All @@ -171,6 +182,67 @@ def forward_post_core_attn(layer: TransformerLayer, args: Dict[str, Any]):
return args


def forward_post_then_pre_core_attn_cuda_graph(layer: TransformerLayer, args: Dict[str, Any]):
log_memory_usage(f"(L{layer.layer_number}) forward_post_then_pre_core_attn:(start)")
assert args["context"] is None and args["context_mask"] is None, "not supported in cudagraph"
forward_post_core_attn_comm(layer, args)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Bug: Previous layer not passed to CUDA graph communication function

In forward_post_then_pre_core_attn_cuda_graph, the function forward_post_core_attn_comm is called with layer (the current layer), but it should use the previous layer for processing the previous layer's attention output. The tick_nonca_compute_cuda_graph function has access to prev_layer but doesn't pass it to forward_post_then_pre_core_attn_cuda_graph. In the non-CUDA-graph version tick_nonca_compute, forward_post_core_attn(prev_layer, arg_group) correctly uses prev_layer. This causes layer._post_attn_to_mlp and config values to be read from the wrong layer.

Additional Locations (1)

Fix in Cursor Fix in Web

log_memory_usage(f"(L{layer.layer_number}) forward_post_then_pre_core_attn:(before non core attn)")
query, key, value, residual, attn_mask_type = layer._forward_pre_attn_cuda_graph(
(args.pop("core_attn_out"), args.pop("residual")),
args["rotary_pos_emb"],
args["rotary_pos_cos"],
args["rotary_pos_sin"],
args["mlp_packed_seq_params"],
args["sequence_len_offset"],
)
log_memory_usage(f"(L{layer.layer_number}) forward_post_then_pre_core_attn:(after non core attn)")
args["query"] = query
args["key"] = key
args["value"] = value
args["residual"] = residual
args["attn_mask_type"] = attn_mask_type
forward_pre_core_attn_comm(layer, args)
log_memory_usage(f"(L{layer.layer_number}) forward_pre_then_pre_core_attn:(return)")
return args


def forward_pre_core_attn_cuda_graph(layer: TransformerLayer, args: Dict[str, Any]):
log_memory_usage(f"(L{layer.layer_number}) forward_pre_core_attn:(start)")
log_memory_usage(f"(L{layer.layer_number}) forward_pre_core_attn:(before pre core attn)")
hidden_states = args.pop("hidden_states")
query, key, value, residual, attn_mask_type = layer._forward_pre_attn_cuda_graph(
(hidden_states,),
args["rotary_pos_emb"],
args["rotary_pos_cos"],
args["rotary_pos_sin"],
args["mlp_packed_seq_params"],
args["sequence_len_offset"],
)
log_memory_usage(f"(L{layer.layer_number}) forward_pre_core_attn:(after pre core attn)")

args["query"] = query
args["key"] = key
args["value"] = value
args["residual"] = residual
args["attn_mask_type"] = attn_mask_type
forward_pre_core_attn_comm(layer, args)
log_memory_usage(f"(L{layer.layer_number}) forward_pre_core_attn:(return)")
return args


def forward_post_core_attn_cuda_graph(layer: TransformerLayer, args: Dict[str, Any]):
forward_post_core_attn_comm(layer, args)
mlp_output = layer._forward_post_attn_cuda_graph(
args.pop("core_attn_out"),
args.pop("residual"),
args["context"],
args["context_mask"],
)
args["hidden_states"] = mlp_output
log_memory_usage(f"(L{layer.layer_number}) forward_post_core_attn:(end)")
return args


def tick_sync(compute_stream, comm_stream, arg_group_0, keys_0, arg_group_1, keys_1,
layer_info="unknown", operation_info="unknown"):
log_memory_usage(f"(L?) tick_sync:(start)")
Expand Down Expand Up @@ -213,6 +285,15 @@ def tick_nonca_compute(
arg_group = forward_pre_core_attn(layer, arg_group)
return arg_group

def tick_nonca_compute_cuda_graph(
layer: TransformerLayer, prev_layer: Optional[TransformerLayer],
arg_group: Dict[str, Any], is_last_layer_post_attn: bool
):
if is_last_layer_post_attn:
return forward_post_core_attn_cuda_graph(layer, arg_group)
if prev_layer is None:
return forward_pre_core_attn(layer, arg_group)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Bug: Wrong function called for first layer in CUDA graph mode

In tick_nonca_compute_cuda_graph, when prev_layer is None (first layer case), it calls forward_pre_core_attn (non-CUDA-graph version) instead of forward_pre_core_attn_cuda_graph. This defeats the purpose of using CUDA graphs and creates inconsistency where the first layer uses the non-graphed path while other layers use the graphed path.

Fix in Cursor Fix in Web

return forward_post_then_pre_core_attn_cuda_graph(layer, arg_group)

# ========== Tick Operations Timing Collection ==========
# TODO: (Refactor) Move this to somewhere else for sole measurement purpose only.
Expand Down
Loading