Skip to content

Commit 1e93933

Browse files
committed
add dsv32
1 parent 51c197c commit 1e93933

5 files changed

Lines changed: 526 additions & 8 deletions

File tree

‎torchtitan/models/deepseek_v3/__init__.py‎

Lines changed: 127 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,9 @@
1515
from torchtitan.distributed.pipeline_parallel import pipeline_llm
1616
from torchtitan.models.common import (
1717
ComplexRoPE,
18+
CosSinRoPE,
1819
Embedding,
20+
LayerNorm,
1921
Linear,
2022
RMSNorm,
2123
RoPE,
@@ -33,7 +35,13 @@
3335
from torchtitan.protocols.model import ModelConfigConverter
3436
from torchtitan.protocols.model_spec import ModelSpec
3537

36-
from .model import Attention, DeepSeekV3Model, DeepSeekV3TransformerBlock
38+
from .model import (
39+
Attention,
40+
DeepSeekV3Model,
41+
DeepSeekV3TransformerBlock,
42+
DeepSeekSparseAttention,
43+
Indexer,
44+
)
3745
from .parallelize import parallelize_deepseekv3
3846
from .state_dict_adapter import DeepSeekV3StateDictAdapter
3947

@@ -88,14 +96,50 @@ def _make_dsv3_attn_config(
8896
mscale: float = 1.0,
8997
attn_backend: str,
9098
rope: RoPE.Config,
99+
index_n_heads: int | None = None,
100+
index_head_dim: int | None = None,
101+
index_topk: int | None = None,
91102
) -> Attention.Config:
92103
"""Build a fully-specified DeepSeek V3 MLA Attention.Config.
93104
94105
All Linear and RMSNorm sub-configs have their dimensional fields set.
95106
When q_lora_rank == 0, sets wq (not wq_a/wq_b).
96107
When q_lora_rank > 0, sets wq_a/wq_b (not wq).
108+
When index_* kwargs are provided, also builds the Lightning Indexer
109+
sub-modules and wraps the inner attention with DeepSeekSparseAttention.
97110
"""
98111
inner_attention = get_attention_config(attn_backend)
112+
indexer = None
113+
if index_n_heads is not None:
114+
assert index_head_dim is not None and index_topk is not None
115+
indexer = Indexer.Config(
116+
dim=dim, q_lora_rank=q_lora_rank,
117+
index_n_heads=index_n_heads, index_head_dim=index_head_dim,
118+
rope_head_dim=qk_rope_head_dim, index_topk=index_topk,
119+
wq_b=Linear.Config(
120+
in_features=q_lora_rank,
121+
out_features=index_n_heads * index_head_dim,
122+
param_init=_LINEAR_INIT,
123+
),
124+
wk=Linear.Config(
125+
in_features=dim, out_features=index_head_dim,
126+
param_init=_LINEAR_INIT,
127+
),
128+
k_norm=LayerNorm.Config(normalized_shape=index_head_dim),
129+
weights_proj=Linear.Config(
130+
in_features=dim, out_features=index_n_heads,
131+
param_init={"weight": partial(nn.init.normal_, std=1.0)},
132+
),
133+
rope=CosSinRoPE.Config(
134+
dim=qk_rope_head_dim, max_seq_len=rope.max_seq_len,
135+
theta=rope.theta, scaling=rope.scaling,
136+
rope_factor=rope.rope_factor,
137+
beta_fast=rope.beta_fast, beta_slow=rope.beta_slow,
138+
original_seq_len=rope.original_seq_len,
139+
),
140+
)
141+
inner_attention = DeepSeekSparseAttention.Config(index_topk=index_topk)
142+
99143
qk_head_dim = qk_nope_head_dim + qk_rope_head_dim
100144

101145
if q_lora_rank == 0:
@@ -154,6 +198,7 @@ def _make_dsv3_attn_config(
154198
),
155199
inner_attention=inner_attention,
156200
rope=dataclasses.replace(rope),
201+
indexer=indexer,
157202
)
158203

159204

@@ -183,6 +228,9 @@ def _build_dsv3_layers(
183228
moe_comm_backend: str,
184229
non_blocking_capacity_factor: float | None,
185230
rope: RoPE.Config,
231+
index_n_heads: int | None = None,
232+
index_head_dim: int | None = None,
233+
index_topk: int | None = None,
186234
) -> list[TransformerBlock.Config]:
187235
"""Build the list of per-layer TransformerBlock configs.
188236
@@ -206,6 +254,9 @@ def _build_dsv3_layers(
206254
mscale=mscale,
207255
attn_backend=attn_backend,
208256
rope=rope,
257+
index_n_heads=index_n_heads,
258+
index_head_dim=index_head_dim,
259+
index_topk=index_topk,
209260
)
210261

211262
if layer_id < n_dense_layers:
@@ -523,8 +574,83 @@ def _671b(
523574
)
524575

525576

577+
def _debugmodel_v32(
578+
attn_backend: str,
579+
moe_comm_backend: str,
580+
non_blocking_capacity_factor: float | None = None,
581+
) -> DeepSeekV3Model.Config:
582+
dim = 256
583+
n_layers = 6
584+
vocab_size = 2048
585+
n_heads = 16
586+
q_lora_rank = 64
587+
kv_lora_rank = 128
588+
qk_nope_head_dim = 64
589+
qk_rope_head_dim = 64
590+
v_head_dim = 64
591+
index_n_heads = 4
592+
index_head_dim = 128
593+
index_topk = 32
594+
moe_hidden_dim = 256
595+
num_shared_experts = 2
596+
dense_hidden_dim = 1024
597+
num_experts = 8
598+
n_dense_layers = 1
599+
600+
layers = _build_dsv3_layers(
601+
n_layers=n_layers,
602+
n_dense_layers=n_dense_layers,
603+
dim=dim,
604+
n_heads=n_heads,
605+
q_lora_rank=q_lora_rank,
606+
kv_lora_rank=kv_lora_rank,
607+
qk_nope_head_dim=qk_nope_head_dim,
608+
qk_rope_head_dim=qk_rope_head_dim,
609+
v_head_dim=v_head_dim,
610+
mscale=0.70,
611+
dense_hidden_dim=dense_hidden_dim,
612+
moe_hidden_dim=moe_hidden_dim,
613+
num_experts=num_experts,
614+
num_shared_experts=num_shared_experts,
615+
router_top_k=3,
616+
router_score_func="softmax",
617+
aux_loss_coeff=1e-4,
618+
attn_backend=attn_backend,
619+
moe_comm_backend=moe_comm_backend,
620+
non_blocking_capacity_factor=non_blocking_capacity_factor,
621+
rope=ComplexRoPE.Config(
622+
dim=qk_rope_head_dim,
623+
max_seq_len=4096 * 4,
624+
theta=10000.0,
625+
scaling="yarn",
626+
rope_factor=40.0,
627+
beta_fast=32.0,
628+
beta_slow=1.0,
629+
original_seq_len=4096,
630+
),
631+
index_n_heads=index_n_heads,
632+
index_head_dim=index_head_dim,
633+
index_topk=index_topk,
634+
)
635+
return DeepSeekV3Model.Config(
636+
vocab_size=vocab_size,
637+
dim=dim,
638+
tok_embeddings=Embedding.Config(
639+
num_embeddings=vocab_size, embedding_dim=dim, param_init=_EMBEDDING_INIT
640+
),
641+
norm=RMSNorm.Config(normalized_shape=dim, param_init=_NORM_INIT),
642+
lm_head=Linear.Config(
643+
in_features=dim,
644+
out_features=vocab_size,
645+
param_init=_output_linear_init(dim),
646+
),
647+
layers=layers,
648+
)
649+
650+
526651
deepseekv3_configs = {
527652
"debugmodel": _debugmodel,
653+
"debugmodel_v32": _debugmodel_v32,
528654
"16B": _16b,
529655
"236B": _236b,
530656
"671B": _671b,

‎torchtitan/models/deepseek_v3/config_registry.py‎

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,41 @@ def deepseek_v3_debugmodel_hybridep() -> Trainer.Config:
7373
return config
7474

7575

76+
def deepseek_v3_debugmodel_v32() -> Trainer.Config:
77+
model_spec = model_registry("debugmodel_v32")
78+
return Trainer.Config(
79+
loss=ChunkedLossWrapper.Config(
80+
loss_fn=CrossEntropyLoss.Config(
81+
global_vocab_size=decoder_vocab_size(model_spec),
82+
),
83+
),
84+
hf_assets_path="./tests/assets/tokenizer",
85+
metrics=MetricsProcessor.Config(log_freq=1),
86+
model_spec=model_spec,
87+
dataloader=HuggingFaceTextDataLoader.Config(dataset="c4_test"),
88+
optimizer=default_adamw(lr=8e-4),
89+
lr_scheduler=LRSchedulersContainer.Config(
90+
warmup_steps=2,
91+
decay_ratio=0.8,
92+
decay_type="linear",
93+
min_lr_factor=0.0,
94+
),
95+
training=TrainingConfig(
96+
local_batch_size=4,
97+
seq_len=2048,
98+
steps=10,
99+
),
100+
parallelism=ParallelismConfig(
101+
expert_parallel_degree=1,
102+
),
103+
checkpoint=CheckpointManager.Config(
104+
interval=10,
105+
last_save_model_only=False,
106+
),
107+
activation_checkpoint=SelectiveAC.Config(),
108+
)
109+
110+
76111
def deepseek_v3_debugmodel_minimal_async_ep() -> Trainer.Config:
77112
config = deepseek_v3_debugmodel()
78113
config.model_spec = model_registry(

0 commit comments

Comments
 (0)