Skip to content

Commit c0b1dc5

Browse files
authored
Deepseek Engram optimization. (#1147)
### PR Category <!-- One of [ Train | Inference | Compress | Serve | RL | Core | Hardware | CICD | Tools | Others ] --> [Train] ### PR Types <!-- One of [ User Experience | New Features | Bug Fixes | Improvements | Performance | Breaking Change| Deprecations | Test Case | Docs | Others ] --> [Improvements] ### PR Description <!-- Describe what you’ve done --> 1. AlltoAll communication when compute multi_head_embedding. 2. Precompute multi_head_embedding. 3. Optional offloading embedding's optimizer states.
1 parent 5d62f33 commit c0b1dc5

12 files changed

Lines changed: 720 additions & 116 deletions

File tree

Lines changed: 149 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,149 @@
1+
# DeepSeek Engram 27B
2+
system:
3+
no_shared_fs: ${experiment.runner.no_shared_fs}
4+
num_workers: 2
5+
tensor_model_parallel_size: 8
6+
expert_model_parallel_size: 8
7+
expert_tensor_parallel_size: 1
8+
context_parallel_size: 1
9+
engram_embedding_parallel_size: 8
10+
sequence_parallel: true
11+
use_distributed_optimizer: true
12+
overlap_grad_reduce: true
13+
overlap_param_gather: true
14+
precision:
15+
bf16: true
16+
attention_softmax_in_fp32: true
17+
accumulate_allreduce_grads_in_fp32: true
18+
logging:
19+
log_interval: 1
20+
tensorboard_log_interval: 1
21+
wandb_project: ${experiment.exp_name}
22+
wandb_exp_name: ${experiment.exp_name}
23+
log_timers_to_tensorboard: true
24+
log_validation_ppl_to_tensorboard: true
25+
log_throughput: true
26+
log_params_norm: true
27+
log_num_zeros_in_grad: true
28+
log_memory_to_tensorboard: true
29+
checkpoint:
30+
save_interval: ${experiment.save_steps}
31+
load: ${experiment.load}
32+
ckpt_format: ${experiment.ckpt_format}
33+
34+
model:
35+
# nsys profile args =================
36+
# profile: true
37+
# profile_step_start: 5
38+
# profile_step_end: 6
39+
# profile_ranks: [0,7] # default [0]
40+
# Note, need to run with nsys profile
41+
42+
# # torch profiler args =================
43+
# profile: true
44+
# use_pytorch_profiler: true
45+
# profile_step_start: 5
46+
# profile_step_end: 6
47+
# profile_ranks: [0] # default [0]
48+
# tensorboard_dir: /workspace/torch_profile
49+
transformer_impl: transformer_engine
50+
num_layers: 30
51+
hidden_size: 2560
52+
num_attention_heads: 32
53+
num_query_groups: 32 # num_key_value_heads
54+
seq_length: 4096
55+
max_position_embeddings: 4096
56+
norm_epsilon: 1e-6
57+
use_rotary_position_embeddings: true
58+
rotary_base: 1000000
59+
swiglu: true
60+
normalization: RMSNorm
61+
qk_layernorm: true
62+
init_method_std: 0.02
63+
attention_dropout: 0.0
64+
hidden_dropout: 0.0
65+
position_embedding_type: rope
66+
untie_embeddings_and_output_weights: true
67+
no_position_embedding: true
68+
no_rope_fusion: true
69+
disable_bias_linear: true
70+
71+
# mla args ==================
72+
multi_latent_attention: true
73+
q_lora_rank: 768
74+
kv_lora_rank: 512
75+
qk_head_dim: 128
76+
qk_pos_emb_head_dim: 64
77+
v_head_dim: 128
78+
79+
# moe args ===================
80+
ffn_hidden_size: 12288
81+
moe_ffn_hidden_size: 1536
82+
moe_grouped_gemm: true
83+
moe_shared_expert_intermediate_size: 3072
84+
num_experts: 56
85+
moe_router_load_balancing_type: "seq_aux_loss"
86+
moe_router_score_function: sigmoid
87+
moe_router_enable_expert_bias: true
88+
moe_router_bias_update_rate: 0.001
89+
moe_aux_loss_coeff: 0.02
90+
moe_layer_freq: "[0]+[1]*29"
91+
# node limited routing
92+
moe_router_num_groups: 1
93+
moe_router_group_topk: 1
94+
moe_router_topk: 6
95+
moe_router_topk_scaling_factor: 2.446
96+
moe_token_dispatcher_type: "alltoall"
97+
# overlap_moe_expert_parallel_comm: true # Optional.
98+
99+
# mtp args ====================
100+
# mtp_num_layers: 1
101+
# mtp_loss_scaling_factor: 0.3
102+
103+
# engram args =================
104+
use_engram: true
105+
engram_tokenizer_name_or_path: /workspace/qwentokenizer
106+
engram_vocab_size: [1131200, 1131200]
107+
max_ngram_size: 3
108+
n_embed_per_ngram: 1280
109+
n_head_per_ngram: 8
110+
engram_layer_ids: [2, 15]
111+
engram_pad_id: 0
112+
engram_seed: 0
113+
engram_kernel_size: 4
114+
engram_hc_mult: 1
115+
engram_embedding_parallel_method: alltoall # alltoall, allreduce, offload
116+
117+
# training
118+
seed: ${experiment.seed}
119+
finetune: false
120+
micro_batch_size: 2
121+
global_batch_size: 2048
122+
eval_iters: 0
123+
train_iters: 20
124+
125+
optimizer:
126+
clip_grad: 1.0
127+
weight_decay: 0.1
128+
adam_beta1: 0.9
129+
adam_beta2: 0.95
130+
lr_scheduler:
131+
lr: 3.0e-3
132+
min_lr: 3.0e-4
133+
lr_warmup_fraction: 0.01
134+
lr_decay_style: WSD
135+
lr_wsd_decay_style: cosine
136+
lr_wsd_decay_iters: 10
137+
138+
data:
139+
reset_position_ids: True
140+
reset_attention_mask: True
141+
data_path: /workspace/data/enron_emails_demo_text_document_qwen
142+
split: 1
143+
no_mmap_bin_files: true
144+
tokenizer:
145+
legacy_tokenizer: true
146+
tokenizer_type: QwenTokenizerFS
147+
tokenizer_path: /workspace/qwentokenizer
148+
vocab_size: 151851
149+
make_vocab_size_divisible_by: 64
Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,51 @@
1+
defaults:
2+
- _self_
3+
- train: engram
4+
5+
experiment:
6+
exp_name: DeepSeek-Engram
7+
seed: 42
8+
save_steps: 100
9+
load: null
10+
exp_dir: outputs/${experiment.exp_name}
11+
ckpt_format: torch
12+
# ckpt_format: fsdp_dtensor # Just for Megatron FSDP.
13+
task:
14+
type: train
15+
backend: megatron
16+
entrypoint: flagscale/train/megatron/train_engram.py
17+
runner:
18+
# 单机
19+
# per_node_task: false
20+
# no_shared_fs: false
21+
# rdzv_backend: static
22+
# hostfile: null
23+
# ssh_port: 10710
24+
# 多机
25+
per_node_task: false
26+
no_shared_fs: false
27+
backend: torchrun
28+
nnodes: 3
29+
nproc_per_node: 8
30+
hostfile: hostfile # Select an available hostfile. Like ip_1 slosts=8\nip_2 slost=8...
31+
master_port: 10720 # Select an available port.
32+
ssh_port: 10710 # Select an available port.
33+
master_addr: <master_ip>
34+
rdzv_backend: static
35+
cmds:
36+
before_start: ulimit -n 1048576 && source /root/miniconda3/bin/activate /root/miniconda3/envs/flagscale-train
37+
envs:
38+
LOGLEVEL: "INFO"
39+
CUDA_VISIBLE_DEVICES: "0,1,2,3,4,5,6,7"
40+
CUDA_DEVICE_MAX_CONNECTIONS: 1
41+
NCCL_IB_HCA: "IB interface" # Select correct IB interface.
42+
NCCL_SOCKET_IFNAME: "IP interface" # Select correct interface.
43+
NCCL_IB_DISABLE: 0
44+
NCCL_DEBUG: "WARN"
45+
NCCL_IB_GID_INDEX: 3
46+
47+
action: run
48+
49+
hydra:
50+
run:
51+
dir: ${experiment.exp_dir}/hydra

flagscale/models/megatron/engram/engram.py

Lines changed: 47 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,9 @@
1414
from .short_conv import ShortConv
1515

1616

17+
## Megatron
18+
from megatron.core.transformer.utils import sharded_state_dict_default
19+
1720
class Engram(nn.Module):
1821
def __init__(self, engram_cfg: EngramConfig, layer_id):
1922
super().__init__()
@@ -34,13 +37,17 @@ def __init__(self, engram_cfg: EngramConfig, layer_id):
3437
pad_id=engram_cfg.engram_pad_id,
3538
seed=engram_cfg.engram_seed,
3639
)
37-
self.multi_head_embedding = MultiHeadEmbedding(
40+
self.memory = MultiHeadEmbedding(
3841
engram_cfg,
3942
list_of_N=[
4043
x for y in global_hash_mapping.vocab_size_across_layers[self.layer_id] for x in y
4144
],
4245
D=engram_cfg.n_embed_per_ngram // engram_cfg.n_head_per_ngram,
4346
)
47+
self.embedding_cache = None # Cache for pre-computed embeddings
48+
self.embedding_stream = None # Stream for pre-computing embeddings
49+
if torch.cuda.is_available():
50+
self.embedding_stream = torch.cuda.Stream()
4451
self.short_conv = ShortConv(
4552
hidden_size=self.backbone_config.hidden_size,
4653
kernel_size=engram_cfg.engram_kernel_size,
@@ -81,8 +88,14 @@ def forward(self, hidden_states, hash_input_ids):
8188
# [B, L, N_GRAM * N_HEADS_PER_GRAM]
8289
# fake hyper-connection
8390
hidden_states = hidden_states.unsqueeze(2)
84-
85-
embeddings = self.multi_head_embedding(hash_input_ids).flatten(start_dim=-2)
91+
if self.embedding_cache is not None:
92+
embeddings, embedding_event = self.embedding_cache
93+
if embedding_event is not None:
94+
torch.cuda.current_stream().wait_event(embedding_event) # Ensure pre-computed embeddings are ready
95+
self.embedding_cache = None # Clear cache after use
96+
del embedding_event # Free the event
97+
else:
98+
embeddings = self.memory(hash_input_ids).flatten(start_dim=-2)
8699
# [L/tp_size, B, N_GRAM * N_HEADS_PER_GRAM, N_EMBED_PER_GRAM // N_HEADS_PER_GRAM]
87100
# [L/tp_size, B, N_GRAM * N_EMBED_PER_NGRAM]
88101

@@ -120,3 +133,34 @@ def forward(self, hidden_states, hash_input_ids):
120133
output = output.squeeze(2)
121134

122135
return output
136+
137+
def pre_compute_embedding(self, input_ids: torch.Tensor):
138+
"""
139+
Pre-compute the multi-head embedding for the given input IDs.
140+
This can be called before the forward pass to warm up the embedding cache.
141+
"""
142+
assert input_ids is not None, "Input ids can not be None for EngramModel"
143+
self.embedding_stream.synchronize() # Ensure previous computations on the stream are finished
144+
with torch.cuda.stream(self.embedding_stream):
145+
embedding_result = self.memory(input_ids).flatten(start_dim=-2)
146+
embedding_event = torch.cuda.Event()
147+
embedding_event.record(self.embedding_stream)
148+
self.embedding_cache = (embedding_result, embedding_event)
149+
150+
def sharded_state_dict(
151+
self, prefix: str = "", sharded_offsets: tuple = (), metadata: dict | None = None
152+
):
153+
sharded_dict = {}
154+
memory_prefix = f"{prefix}memory."
155+
sharded_dict.update(self.memory.sharded_state_dict(memory_prefix, sharded_offsets, metadata))
156+
conv_prefix = f"{prefix}short_conv."
157+
sharded_dict.update(sharded_state_dict_default(self.short_conv, conv_prefix, sharded_offsets, metadata))
158+
value_proj_prefix = f"{prefix}value_proj."
159+
sharded_dict.update(sharded_state_dict_default(self.value_proj, value_proj_prefix, sharded_offsets, metadata))
160+
key_projs_prefix = f"{prefix}key_projs."
161+
sharded_dict.update(sharded_state_dict_default(self.key_projs, key_projs_prefix, sharded_offsets, metadata))
162+
norm1_prefix = f"{prefix}norm1."
163+
sharded_dict.update(sharded_state_dict_default(self.norm1, norm1_prefix, sharded_offsets, metadata))
164+
norm2_prefix = f"{prefix}norm2."
165+
sharded_dict.update(sharded_state_dict_default(self.norm2, norm2_prefix, sharded_offsets, metadata))
166+
return sharded_dict

flagscale/models/megatron/engram/engram_config.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,3 +17,6 @@ class EngramConfig(MLATransformerConfig):
1717
engram_seed: int = 0
1818
engram_kernel_size: int = 1
1919
engram_hc_mult: int = 1
20+
engram_embedding_parallel_size: int | None = 1
21+
engram_embedding_parallel_method: str = "alltoall"
22+
engram_offload_embedding_optimizer_states: bool = False

flagscale/models/megatron/engram/engram_model.py

Lines changed: 68 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
# ruff: noqa: RUF013
22
## built-in
3+
from typing import Optional
4+
35
import torch
46
from torch import Tensor
57

@@ -27,26 +29,41 @@ def __init__(self, hash_mapping, input_ids, hash_stream=None):
2729
self.input_ids = input_ids
2830
self.hash_stream = hash_stream
2931
self._result = None
30-
self._computation_started = False
31-
32-
# torch.cuda.nvtx.range_push("LazyHashInputIds hash")
33-
# Start async computation immediately if stream is available
32+
self._is_async_pending = False
33+
# Async
3434
if self.hash_stream is not None:
35+
# self.hash_stream.wait_stream(torch.cuda.current_stream())
3536
with torch.cuda.stream(self.hash_stream):
3637
self._result = self.hash_mapping.hash(self.input_ids)
37-
self._computation_started = True
38-
# torch.cuda.nvtx.range_pop()
38+
self._is_async_pending = True
39+
# record result to use across stream
40+
self._record_current_stream()
3941

40-
def __getitem__(self, key):
41-
"""Access hash result, synchronizing if necessary."""
42+
def _record_current_stream(self):
43+
"""Helper to record current stream on all result tensors"""
4244
if self._result is None:
43-
if self.hash_stream is not None and self._computation_started:
44-
# Wait for async computation to complete
45-
torch.cuda.current_stream().wait_stream(self.hash_stream)
46-
self._computation_started = False # Mark as synchronized
47-
else:
48-
# Compute synchronously if no stream or computation not started
49-
self._result = self.hash_mapping.hash(self.input_ids)
45+
return
46+
current_stream = torch.cuda.current_stream()
47+
if isinstance(self._result, dict):
48+
for t in self._result.values():
49+
if isinstance(t, torch.Tensor):
50+
t.record_stream(current_stream)
51+
elif isinstance(self._result, torch.Tensor):
52+
self._result.record_stream(current_stream)
53+
54+
def __getitem__(self, key):
55+
# Case 1: Async compute -> wait
56+
if self._is_async_pending:
57+
torch.cuda.current_stream().wait_stream(self.hash_stream)
58+
self._is_async_pending = False # Async finish
59+
self._record_current_stream()
60+
61+
# Case 2: Sync but no compute -> start compute
62+
elif self._result is None:
63+
self._result = self.hash_mapping.hash(self.input_ids)
64+
65+
# Case 3: Async or sync compute is finished.
66+
# print(f"[rank{torch.distributed.get_rank()}]: LazyHashInputIds result = {self._result}")
5067
return self._result[key]
5168

5269
def get(self, key, default=None):
@@ -171,7 +188,40 @@ def forward(
171188
inference_context=inference_context,
172189
)
173190

174-
def sharded_state_dict(
175-
self, prefix: str = "", sharded_offsets: tuple = (), metadata: dict | None = None
191+
def build_schedule_plan(
192+
self,
193+
input_ids: Tensor,
194+
position_ids: Tensor,
195+
attention_mask: Tensor,
196+
decoder_input: Tensor = None,
197+
labels: Tensor = None,
198+
inference_context: BaseInferenceContext = None,
199+
packed_seq_params: PackedSeqParams = None,
200+
extra_block_kwargs: dict = None,
201+
runtime_gather_output: Optional[bool] = None,
202+
inference_params: Optional[BaseInferenceContext] = None,
203+
loss_mask: Optional[Tensor] = None,
176204
):
177-
raise NotImplementedError("Sharded state dict is not supported for EngramModel")
205+
"""
206+
Adaptation of overlap_moe_expert_parallel_comm.
207+
"""
208+
# Precompute the engram_hash_iput_ids, it will be used to create a TransformerChunkSchedulePlan.
209+
engram_hash_input_ids = LazyHashInputIds(
210+
hash_mapping=self.engram_hash,
211+
input_ids=input_ids,
212+
hash_stream=self._hash_stream,
213+
)
214+
if extra_block_kwargs is None:
215+
extra_block_kwargs = {
216+
"engram_hash_input_ids": engram_hash_input_ids,
217+
}
218+
return super().build_schedule_plan(
219+
input_ids,
220+
position_ids,
221+
attention_mask,
222+
decoder_input,
223+
labels=labels,
224+
loss_mask=loss_mask,
225+
extra_block_kwargs=extra_block_kwargs
226+
)
227+

0 commit comments

Comments
 (0)