Skip to content

Commit b9aee47

Browse files
Merge pull request #4367 from Shuwen-Fang:deprecate_tp_transpose
PiperOrigin-RevId: 943625535
2 parents cdebf0f + faaa424 commit b9aee47

15 files changed

Lines changed: 145 additions & 1813 deletions

File tree

src/maxtext/configs/base.yml

Lines changed: 18 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -492,7 +492,7 @@ compile_xla_flags: "" # Compiler options e.g. compile_xla_flags="--xla_tpu_num_s
492492
# Parallelism
493493
shard_mode: "auto" # can be either auto or explicit
494494
custom_mesh_and_rule: "" # replace default mesh and logical rule by specifying yml name under config/mesh_and_rule/.
495-
mesh_axes: ['diloco', 'data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'context_autoregressive', 'tensor', 'tensor_transpose', 'tensor_sequence', 'expert', 'autoregressive']
495+
mesh_axes: ['diloco', 'data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'context_autoregressive', 'tensor', 'tensor_sequence', 'expert', 'autoregressive']
496496
logical_axis_rules: [
497497
['circular_repeats', []],
498498
# ==========================================
@@ -501,40 +501,36 @@ logical_axis_rules: [
501501
# Vocab Activations
502502
['activation_embed_and_logits_batch', ['data', 'stage', 'fsdp', 'fsdp_transpose', 'expert']],
503503
['activation_embed_and_logits_batch_sequence', ['data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'expert']],
504-
['activation_vocab', ['tensor', 'tensor_transpose', 'tensor_sequence']],
505-
['activation_vocab', ['tensor', 'tensor_transpose']],
504+
['activation_vocab', ['tensor', 'tensor_sequence']],
505+
['activation_vocab', ['tensor']],
506506
['activation_vocab', 'tensor_sequence'],
507507
# Vocab Weights
508-
['vocab', ['tensor', 'tensor_transpose', 'tensor_sequence', 'autoregressive']],
508+
['vocab', ['tensor', 'tensor_sequence', 'autoregressive']],
509509
['embed_vocab', ['fsdp', 'fsdp_transpose', 'context', 'expert']],
510510
# ==========================================
511511
# Attention
512512
# ==========================================
513513
# Attention Activations
514514
['activation_batch_attn', ['data', 'fsdp', 'fsdp_transpose', 'expert']],
515-
['activation_heads', ['tensor', 'tensor_transpose', 'tensor_sequence', 'autoregressive']],
516-
['activation_kv_heads', ['tensor', 'tensor_transpose', 'tensor_sequence']],
515+
['activation_heads', ['tensor', 'tensor_sequence', 'autoregressive']],
516+
['activation_kv_heads', ['tensor', 'tensor_sequence']],
517517
['activation_length_attn', ['context']],
518518
['activation_q_length', ['context']],
519519
['activation_kv_length', []],
520-
['activation_embed_attn', ['tensor', 'tensor_transpose']],
521-
['activation_kv', ['tensor', 'tensor_transpose', 'tensor_sequence']],
520+
['activation_embed_attn', ['tensor']],
521+
['activation_kv', ['tensor', 'tensor_sequence']],
522522
['activation_kv_batch', ['data', 'fsdp', 'fsdp_transpose', 'expert']],
523-
['activation_kv_head_dim', ['tensor', 'tensor_transpose', 'tensor_sequence']],
523+
['activation_kv_head_dim', ['tensor', 'tensor_sequence']],
524524
# Attention Weights
525-
['heads', ['tensor', 'tensor_transpose', 'tensor_sequence', 'autoregressive']],
526-
['q_heads', ['tensor', 'tensor_transpose', 'tensor_sequence', 'autoregressive']],
527-
['kv_heads', ['tensor', 'tensor_transpose', 'tensor_sequence', 'autoregressive']],
525+
['heads', ['tensor', 'tensor_sequence', 'autoregressive']],
526+
['q_heads', ['tensor', 'tensor_sequence', 'autoregressive']],
527+
['kv_heads', ['tensor', 'tensor_sequence', 'autoregressive']],
528528
['qkv', []],
529529
['kv', []],
530530
['kv_head_dim', []],
531-
['q_lora', ['fsdp', 'fsdp_transpose', 'context', 'tensor_transpose', 'expert']],
532-
['q_lora', ['fsdp', 'context', 'tensor_transpose', 'expert']],
533531
['q_lora', ['fsdp', 'fsdp_transpose', 'context', 'expert']],
534532
['q_lora', ['fsdp', 'context', 'expert']],
535533
["q_lora_up_proj", []],
536-
['kv_lora', ['fsdp', 'fsdp_transpose', 'context', 'tensor_transpose', 'expert']],
537-
['kv_lora', ['fsdp', 'context', 'tensor_transpose', 'expert']],
538534
['kv_lora', ['fsdp', 'fsdp_transpose', 'context', 'expert']],
539535
['kv_lora', ['fsdp', 'context', 'expert']],
540536
["kv_lora_up_proj", []],
@@ -545,37 +541,33 @@ logical_axis_rules: [
545541
['activation_batch_moe', ['data', 'fsdp', 'fsdp_transpose', 'expert']],
546542
['activation_length_moe', ['context']],
547543
['activation_norm_length_moe', ['tensor_sequence', 'context']],
548-
['activation_embed_moe', ['tensor', 'tensor_transpose']],
549-
['activation_mlp_moe', ['tensor', 'tensor_transpose', 'tensor_sequence']],
544+
['activation_embed_moe', ['tensor']],
545+
['activation_mlp_moe', ['tensor', 'tensor_sequence']],
550546
['activation_exp', ['expert']],
551547
# MoE Weights
552548
['exp', 'expert'],
553549
['mlp_moe', ['fsdp_transpose', 'tensor', 'tensor_sequence', 'autoregressive']],
554-
['embed_moe', ['fsdp', 'fsdp_transpose', 'tensor_transpose', 'context']],
555-
['embed_moe', ['fsdp', 'tensor_transpose', 'context']],
556550
['embed_moe', ['fsdp', 'fsdp_transpose', 'context']],
557551
['embed_moe', ['fsdp', 'context']],
558552
# ==========================================
559553
# Standard MLP / Dense Layers / Model Structure
560554
# ==========================================
561555
# Dense Activations
562-
['activation_mlp', ['tensor', 'tensor_transpose', 'tensor_sequence']],
556+
['activation_mlp', ['tensor', 'tensor_sequence']],
563557
# Note activation batch and length also get used in vocab
564558
['activation_batch', ['data', 'fsdp', 'fsdp_transpose', 'expert']],
565559
['activation_length', ['context']],
566560
['activation_norm_length', ['tensor_sequence', 'context']],
567-
['activation_embed', ['tensor', 'tensor_transpose']],
561+
['activation_embed', ['tensor']],
568562
['activation_stage', 'stage'],
569563
# General Weights
570564
['mlp', ['fsdp_transpose', 'tensor', 'tensor_sequence', 'autoregressive']],
571565
# GDN (linear-attention) projections shard like 'mlp' during training; the
572566
# vLLM serving config overrides this to match tpu-inference's ATTN_HEAD order.
573567
['gdn_head', ['fsdp_transpose', 'tensor', 'tensor_sequence', 'autoregressive']],
574-
['embed', ['fsdp', 'fsdp_transpose', 'tensor_transpose', 'context', 'expert']],
575-
['embed', ['fsdp', 'tensor_transpose', 'context', 'expert']],
576568
['embed', ['fsdp', 'fsdp_transpose', 'context', 'expert']],
577569
['embed', ['fsdp', 'context', 'expert']],
578-
['norm', ['tensor', 'tensor_transpose']],
570+
['norm', ['tensor']],
579571
['layers', 'stage'],
580572
['diloco', 'diloco'],
581573
['engram_dim', ['tensor']],
@@ -590,7 +582,6 @@ logical_axis_rules: [
590582
['activation_prefill_kv_batch', ['data', 'fsdp', 'fsdp_transpose', 'expert']],
591583
['decode_batch', ['data', 'fsdp', 'fsdp_transpose', 'expert']],
592584
['decode_length', []],
593-
['cache_heads', ['autoregressive', 'tensor', 'tensor_transpose', 'tensor_sequence']],
594585
['cache_heads', ['autoregressive', 'tensor', 'tensor_sequence']],
595586
['paged_kv_heads', ['tensor']],
596587
['cache_batch_prefill', []],
@@ -605,11 +596,10 @@ logical_axis_rules: [
605596
# Deprecated / Scheduled for Removal
606597
# ==========================================
607598
['mlp_no_fsdp', ['tensor', 'tensor_sequence', 'autoregressive']],
608-
['embed_tensor_transpose', ['tensor_transpose']],
609599
['exp_with_fsdp', 'fsdp'],
610600
]
611601
# Axes used for DCN must be earlier in this list than ICI, see (b/339009148) for details
612-
data_sharding: [['data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'context_autoregressive', 'tensor', 'tensor_transpose', 'tensor_sequence', 'expert', 'autoregressive']]
602+
data_sharding: [['data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'context_autoregressive', 'tensor', 'tensor_sequence', 'expert', 'autoregressive']]
613603
input_data_sharding_logical_axes: ['activation_embed_and_logits_batch', 'activation_norm_length']
614604
# Determines which physical axis plays the role of context parallelism for input data processing and load balancing
615605
# only supports "context" or "expert" (when custom_mesh_and_rule=ep-as-cp)
@@ -630,7 +620,6 @@ dcn_sequence_parallelism: 1 # never recommended
630620
dcn_context_parallelism: 1
631621
dcn_context_autoregressive_parallelism: 1
632622
dcn_tensor_parallelism: 1 # never recommended
633-
dcn_tensor_transpose_parallelism: 1
634623
dcn_tensor_sequence_parallelism: 1 # never recommended
635624
dcn_pipeline_parallelism: 1
636625
dcn_expert_parallelism: 1
@@ -643,7 +632,6 @@ ici_sequence_parallelism: 1
643632
ici_context_parallelism: 1
644633
ici_context_autoregressive_parallelism: 1
645634
ici_tensor_parallelism: 1
646-
ici_tensor_transpose_parallelism: 1
647635
ici_tensor_sequence_parallelism: 1
648636
ici_autoregressive_parallelism: 1
649637
ici_pipeline_parallelism: 1

src/maxtext/configs/decoupled_base_test.yml

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -42,7 +42,6 @@ ici_pipeline_parallelism: 1
4242
ici_expert_parallelism: 1
4343
ici_sequence_parallelism: 1
4444
ici_context_parallelism: 1
45-
ici_tensor_transpose_parallelism: 1
4645
ici_tensor_sequence_parallelism: 1
4746
ici_autoregressive_parallelism: 1
4847
ici_fsdp_parallelism: -1
@@ -57,7 +56,6 @@ dcn_pipeline_parallelism: 1
5756
dcn_expert_parallelism: 1
5857
dcn_sequence_parallelism: 1
5958
dcn_context_parallelism: 1
60-
dcn_tensor_transpose_parallelism: 1
6159
dcn_tensor_sequence_parallelism: 1
6260
dcn_autoregressive_parallelism: 1
6361
dcn_fsdp_parallelism: 1

src/maxtext/configs/pyconfig_deprecated.py

Lines changed: 1 addition & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -918,7 +918,6 @@ def create_parallelisms_list(raw_keys):
918918
raw_keys["ici_context_parallelism"],
919919
raw_keys["ici_context_autoregressive_parallelism"],
920920
raw_keys["ici_tensor_parallelism"],
921-
raw_keys["ici_tensor_transpose_parallelism"],
922921
raw_keys["ici_tensor_sequence_parallelism"],
923922
raw_keys["ici_expert_parallelism"],
924923
raw_keys["ici_autoregressive_parallelism"],
@@ -932,7 +931,6 @@ def create_parallelisms_list(raw_keys):
932931
raw_keys["dcn_context_parallelism"],
933932
raw_keys["dcn_context_autoregressive_parallelism"],
934933
raw_keys["dcn_tensor_parallelism"],
935-
raw_keys["dcn_tensor_transpose_parallelism"],
936934
raw_keys["dcn_tensor_sequence_parallelism"],
937935
raw_keys["dcn_expert_parallelism"],
938936
raw_keys["dcn_autoregressive_parallelism"],
@@ -1029,7 +1027,6 @@ def pipeline_first_axis(raw_keys):
10291027
raw_keys["ici_context_parallelism"],
10301028
raw_keys["ici_context_autoregressive_parallelism"],
10311029
raw_keys["ici_tensor_parallelism"],
1032-
raw_keys["ici_tensor_transpose_parallelism"],
10331030
raw_keys["ici_tensor_sequence_parallelism"],
10341031
raw_keys["ici_expert_parallelism"],
10351032
raw_keys["ici_autoregressive_parallelism"],
@@ -1043,7 +1040,6 @@ def pipeline_first_axis(raw_keys):
10431040
raw_keys["dcn_context_parallelism"],
10441041
raw_keys["dcn_context_autoregressive_parallelism"],
10451042
raw_keys["dcn_tensor_parallelism"],
1046-
raw_keys["dcn_tensor_transpose_parallelism"],
10471043
raw_keys["dcn_tensor_sequence_parallelism"],
10481044
raw_keys["dcn_expert_parallelism"],
10491045
raw_keys["dcn_autoregressive_parallelism"],
@@ -1057,7 +1053,6 @@ def pipeline_first_axis(raw_keys):
10571053
"context",
10581054
"context_autoregressive",
10591055
"tensor",
1060-
"tensor_transpose",
10611056
"tensor_sequence",
10621057
"expert",
10631058
"autoregressive",
@@ -1072,7 +1067,6 @@ def pipeline_first_axis(raw_keys):
10721067
"context",
10731068
"context_autoregressive",
10741069
"tensor",
1075-
"tensor_transpose",
10761070
"tensor_sequence",
10771071
"expert",
10781072
"autoregressive",
@@ -1217,9 +1211,7 @@ def validate_shard_expert_on_fsdp(raw_keys):
12171211
if raw_keys["shard_exp_on_fsdp"] and raw_keys["num_experts"] % raw_keys["ici_fsdp_parallelism"] != 0:
12181212
raise ValueError("shard_exp_on_fsdp requires num_experts is divisiable by ici_fsdp_parallelism.")
12191213
if raw_keys["shard_exp_on_fsdp"] and (using_tensor_parallelism(raw_keys) or using_expert_parallelism(raw_keys)):
1220-
raise ValueError(
1221-
"shard_exp_on_fsdp requires ici_expert_parallelism = 1 and ici_tensor_parallelism/ici_tensor_transpose_parallelism = 1."
1222-
)
1214+
raise ValueError("shard_exp_on_fsdp requires ici_expert_parallelism = 1 and ici_tensor_parallelism = 1.")
12231215

12241216

12251217
def validate_ragged_dot(raw_keys):

src/maxtext/configs/types.py

Lines changed: 0 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -963,7 +963,6 @@ class HardwareAndMesh(BaseModel):
963963
"context",
964964
"context_autoregressive",
965965
"tensor",
966-
"tensor_transpose",
967966
"tensor_sequence",
968967
"expert",
969968
"autoregressive",
@@ -1036,7 +1035,6 @@ class DcnParallelism(BaseModel):
10361035
dcn_context_parallelism: int = Field(1, description="DCN axis for context parallelism.")
10371036
dcn_context_autoregressive_parallelism: int = Field(1, description="DCN axis for context autoregressive parallelism.")
10381037
dcn_tensor_parallelism: int = Field(1, description="DCN axis for tensor parallelism (not recommended).")
1039-
dcn_tensor_transpose_parallelism: int = Field(1, description="DCN axis for tensor transpose parallelism.")
10401038
dcn_tensor_sequence_parallelism: int = Field(
10411039
1, description="DCN axis for tensor sequence parallelism (not recommended)."
10421040
)
@@ -1056,7 +1054,6 @@ class IciParallelism(BaseModel):
10561054
ici_context_parallelism: int = Field(1, description="ICI axis for context parallelism.")
10571055
ici_context_autoregressive_parallelism: int = Field(1, description="ICI axis for context autoregressive parallelism.")
10581056
ici_tensor_parallelism: int = Field(1, description="ICI axis for tensor parallelism.")
1059-
ici_tensor_transpose_parallelism: int = Field(1, description="ICI axis for tensor transpose parallelism.")
10601057
ici_tensor_sequence_parallelism: int = Field(1, description="ICI axis for tensor sequence parallelism.")
10611058
ici_autoregressive_parallelism: int = Field(1, description="ICI axis for autoregressive parallelism.")
10621059
ici_pipeline_parallelism: int = Field(1, description="ICI axis for pipeline parallelism.")
@@ -3407,7 +3404,6 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de
34073404
"context": self.ici_context_parallelism,
34083405
"context_autoregressive": self.ici_context_autoregressive_parallelism,
34093406
"tensor": self.ici_tensor_parallelism,
3410-
"tensor_transpose": self.ici_tensor_transpose_parallelism,
34113407
"tensor_sequence": self.ici_tensor_sequence_parallelism,
34123408
"model": self.ici_tensor_parallelism,
34133409
"expert": self.ici_expert_parallelism,
@@ -3427,7 +3423,6 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de
34273423
"context": self.dcn_context_parallelism,
34283424
"context_autoregressive": self.dcn_context_autoregressive_parallelism,
34293425
"tensor": self.dcn_tensor_parallelism,
3430-
"tensor_transpose": self.dcn_tensor_transpose_parallelism,
34313426
"tensor_sequence": self.dcn_tensor_sequence_parallelism,
34323427
"model": self.dcn_tensor_parallelism,
34333428
"expert": self.dcn_expert_parallelism,

0 commit comments

Comments
 (0)