From bdafbc926eb59931d071118d77cf9e4a01ff26eb Mon Sep 17 00:00:00 2001 From: continuousml Date: Mon, 10 Aug 2026 21:12:52 -0700 Subject: [PATCH] Add TPU USP context parallelism --- docs/guides/optimization/sharding.md | 2 + src/maxtext/configs/base.yml | 47 +-- src/maxtext/configs/types.py | 125 ++++++- .../kernels/attention/ulysses_attention.py | 28 +- .../kernels/attention/usp_attention.py | 167 +++++++++ src/maxtext/layers/attention_op.py | 141 ++++++- src/maxtext/utils/maxtext_utils.py | 2 + src/maxtext/utils/sharding.py | 1 + tests/unit/attention_test.py | 160 ++++++++ tests/unit/configs_value_test.py | 122 ++++++ tests/unit/ulysses_attention_test.py | 11 +- tests/unit/usp_attention_test.py | 160 ++++++++ tests/unit/usp_collective_test.py | 170 +++++++++ .../slice_1/rule_default/named_shardings.json | 252 ++++++++++++- .../named_shardings.json | 252 ++++++++++++- .../slice_1/rule_default/named_shardings.json | 348 +++++++++++++++++- .../named_shardings.json | 348 +++++++++++++++++- .../slice_1/rule_default/named_shardings.json | 108 ++++++ 18 files changed, 2339 insertions(+), 105 deletions(-) create mode 100644 src/maxtext/kernels/attention/usp_attention.py create mode 100644 tests/unit/usp_attention_test.py create mode 100644 tests/unit/usp_collective_test.py diff --git a/docs/guides/optimization/sharding.md b/docs/guides/optimization/sharding.md index 5927914995..92db416086 100644 --- a/docs/guides/optimization/sharding.md +++ b/docs/guides/optimization/sharding.md @@ -262,6 +262,8 @@ MaxText supports `context_parallel_strategy=all_gather`, and supports `context_p MaxText also supports `context_parallel_strategy=ulysses` ([DeepSpeed Ulysses](https://arxiv.org/abs/2309.14509)) on the TPU Tokamax Splash path for training. Ulysses exchanges sequence ownership for head ownership by communicating the Q, K, V, and output activations through all-to-all collectives: each device computes ordinary full-sequence attention for its head subset, and the inverse all-to-all restores the sequence sharding on the output. It requires explicit positive context parallelism values, `context_sharding=context`, `attention=flash` with Tokamax Splash, global causal attention, query and KV head counts divisible by the context parallel size including after tensor-parallel head sharding, matching Q and KV head-sharding axes, an unsharded head feature dimension, a divisible sequence length, `dq_reduction_steps` of 0 or 3, `context_parallel_load_balance=false` (each device computes full-sequence attention for its head subset, so the work is already balanced and the causal load-balancing reorder must stay off), and ICI-only context parallelism (`dcn_context_parallelism` must equal 1). It does not support MQA, dropout, QK-Clip statistics, ragged attention, attention sinks, sparse indexer masks, chunked prefill, MoBA, or multimodal attention. +MaxText also supports `context_parallel_strategy=usp` ([USP](https://arxiv.org/abs/2405.07719), Ulysses over ring) on the same TPU Tokamax Splash path for training. This initial support is non-load-balanced. USP factors the context parallelism into a ring dimension on the `context` mesh axis (`ici_context_parallelism`) and a Ulysses dimension on the `context_ulysses` mesh axis (`ici_context_ulysses_parallelism`): the Ulysses all-to-all exchanges sequence ownership for head ownership over the Ulysses axis at each fixed ring position, and the ring kernel then rotates K and V across the ring axis inside each head subset. The strategy is hybrid-only: both dimensions must be greater than one, and the single-dimension endpoints are the existing `ring` and `ulysses` strategies. It shares the Ulysses restrictions (explicit positive ICI-only sizes, `attention=flash` with Tokamax Splash, global causal attention, head counts divisible by the Ulysses size, no MQA, no packing, no load balancing, no dropout, no multi-token prediction, no dKV megacore) and additionally requires `max_target_length` divisible by the total context parallelism and by the ring size squared. + ### CP Arithmetic Intensity The main communications are the same as FSDP (all gather weights and synchronize gradients), with an arithmetic intensity of `local_batch` / `sparsity`. diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index 08d22bae24..72265e5fbe 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -514,7 +514,7 @@ compile_xla_flags: "" # Compiler options e.g. compile_xla_flags="--xla_tpu_num_s # Parallelism shard_mode: "auto" # can be either auto or explicit custom_mesh_and_rule: "" # replace default mesh and logical rule by specifying yml name under config/mesh_and_rule/. -mesh_axes: ['diloco', 'data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'context_autoregressive', 'tensor', 'tensor_sequence', 'expert', 'autoregressive'] +mesh_axes: ['diloco', 'data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'context_ulysses', 'context_autoregressive', 'tensor', 'tensor_sequence', 'expert', 'autoregressive'] logical_axis_rules: [ ['circular_repeats', []], # ========================================== @@ -522,13 +522,13 @@ logical_axis_rules: [ # ========================================== # Vocab Activations ['activation_embed_and_logits_batch', ['data', 'stage', 'fsdp', 'fsdp_transpose', 'expert']], - ['activation_embed_and_logits_batch_sequence', ['data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'expert']], + ['activation_embed_and_logits_batch_sequence', ['data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'context_ulysses', 'expert']], ['activation_vocab', ['tensor', 'tensor_sequence']], ['activation_vocab', ['tensor']], ['activation_vocab', 'tensor_sequence'], # Vocab Weights ['vocab', ['tensor', 'tensor_sequence', 'autoregressive']], - ['embed_vocab', ['fsdp', 'fsdp_transpose', 'context', 'expert']], + ['embed_vocab', ['fsdp', 'fsdp_transpose', 'context', 'context_ulysses', 'expert']], # ========================================== # Attention # ========================================== @@ -536,8 +536,8 @@ logical_axis_rules: [ ['activation_batch_attn', ['data', 'fsdp', 'fsdp_transpose', 'expert']], ['activation_heads', ['tensor', 'tensor_sequence', 'autoregressive']], ['activation_kv_heads', ['tensor', 'tensor_sequence']], - ['activation_length_attn', ['context']], - ['activation_q_length', ['context']], + ['activation_length_attn', ['context', 'context_ulysses']], + ['activation_q_length', ['context', 'context_ulysses']], ['activation_kv_length', []], ['activation_embed_attn', ['tensor']], ['activation_kv', ['tensor', 'tensor_sequence']], @@ -550,27 +550,27 @@ logical_axis_rules: [ ['qkv', []], ['kv', []], ['kv_head_dim', []], - ['q_lora', ['fsdp', 'fsdp_transpose', 'context', 'expert']], - ['q_lora', ['fsdp', 'context', 'expert']], + ['q_lora', ['fsdp', 'fsdp_transpose', 'context', 'context_ulysses', 'expert']], + ['q_lora', ['fsdp', 'context', 'context_ulysses', 'expert']], ["q_lora_up_proj", []], - ['kv_lora', ['fsdp', 'fsdp_transpose', 'context', 'expert']], - ['kv_lora', ['fsdp', 'context', 'expert']], + ['kv_lora', ['fsdp', 'fsdp_transpose', 'context', 'context_ulysses', 'expert']], + ['kv_lora', ['fsdp', 'context', 'context_ulysses', 'expert']], ["kv_lora_up_proj", []], # ========================================== # Mixture of Experts (MoE) # ========================================== # MoE Activations ['activation_batch_moe', ['data', 'fsdp', 'fsdp_transpose', 'expert']], - ['activation_length_moe', ['context']], - ['activation_norm_length_moe', ['tensor_sequence', 'context']], + ['activation_length_moe', ['context', 'context_ulysses']], + ['activation_norm_length_moe', ['tensor_sequence', 'context', 'context_ulysses']], ['activation_embed_moe', ['tensor']], ['activation_mlp_moe', ['tensor', 'tensor_sequence']], ['activation_exp', ['expert']], # MoE Weights ['exp', 'expert'], ['mlp_moe', ['fsdp_transpose', 'tensor', 'tensor_sequence', 'autoregressive']], - ['embed_moe', ['fsdp', 'fsdp_transpose', 'context']], - ['embed_moe', ['fsdp', 'context']], + ['embed_moe', ['fsdp', 'fsdp_transpose', 'context', 'context_ulysses']], + ['embed_moe', ['fsdp', 'context', 'context_ulysses']], # ========================================== # Standard MLP / Dense Layers / Model Structure # ========================================== @@ -578,8 +578,8 @@ logical_axis_rules: [ ['activation_mlp', ['tensor', 'tensor_sequence']], # Note activation batch and length also get used in vocab ['activation_batch', ['data', 'fsdp', 'fsdp_transpose', 'expert']], - ['activation_length', ['context']], - ['activation_norm_length', ['tensor_sequence', 'context']], + ['activation_length', ['context', 'context_ulysses']], + ['activation_norm_length', ['tensor_sequence', 'context', 'context_ulysses']], ['activation_embed', ['tensor']], ['activation_stage', 'stage'], # General Weights @@ -587,8 +587,8 @@ logical_axis_rules: [ # GDN (linear-attention) projections shard like 'mlp' during training; the # vLLM serving config overrides this to match tpu-inference's ATTN_HEAD order. ['gdn_head', ['fsdp_transpose', 'tensor', 'tensor_sequence', 'autoregressive']], - ['embed', ['fsdp', 'fsdp_transpose', 'context', 'expert']], - ['embed', ['fsdp', 'context', 'expert']], + ['embed', ['fsdp', 'fsdp_transpose', 'context', 'context_ulysses', 'expert']], + ['embed', ['fsdp', 'context', 'context_ulysses', 'expert']], ['norm', ['tensor']], ['layers', 'stage'], ['diloco', 'diloco'], @@ -600,8 +600,8 @@ logical_axis_rules: [ # ========================================== # Inference(Prefill, Decode, Cache) # ========================================== - ['prefill_activation_length', ['context']], - ['prefill_activation_norm_length', ['tensor_sequence', 'context']], + ['prefill_activation_length', ['context', 'context_ulysses']], + ['prefill_activation_norm_length', ['tensor_sequence', 'context', 'context_ulysses']], ['activation_prefill_kv_batch', ['data', 'fsdp', 'fsdp_transpose', 'expert']], ['decode_batch', ['data', 'fsdp', 'fsdp_transpose', 'expert']], ['decode_length', []], @@ -622,11 +622,14 @@ logical_axis_rules: [ ['exp_with_fsdp', 'fsdp'], ] # Axes used for DCN must be earlier in this list than ICI, see (b/339009148) for details -data_sharding: [['data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'context_autoregressive', 'tensor', 'tensor_sequence', 'expert', 'autoregressive']] +data_sharding: [['data', 'stage', 'fsdp', 'fsdp_transpose', 'context', 'context_ulysses', 'context_autoregressive', 'tensor', 'tensor_sequence', 'expert', 'autoregressive']] input_data_sharding_logical_axes: ['activation_embed_and_logits_batch', 'activation_norm_length'] # Determines which physical axis plays the role of context parallelism for input data processing and load balancing # only supports "context" or "expert" (when custom_mesh_and_rule=ep-as-cp) context_sharding: "context" +# Physical axis for the Ulysses (all-to-all) dimension of context_parallel_strategy='usp' +# only supports "context_ulysses" +ulysses_context_sharding: "context_ulysses" # Customized mesh and logical rules for evaluation (e.g. "ep-as-cp") custom_mesh_and_rule_for_eval: "" logical_axis_rules_for_eval: [] @@ -644,6 +647,7 @@ dcn_fsdp_parallelism: 1 dcn_fsdp_transpose_parallelism: 1 dcn_sequence_parallelism: 1 # never recommended dcn_context_parallelism: 1 +dcn_context_ulysses_parallelism: 1 dcn_context_autoregressive_parallelism: 1 dcn_tensor_parallelism: 1 # never recommended dcn_tensor_sequence_parallelism: 1 # never recommended @@ -656,6 +660,7 @@ ici_fsdp_parallelism: -1 # recommended ICI axis to be auto-sharded ici_fsdp_transpose_parallelism: 1 ici_sequence_parallelism: 1 ici_context_parallelism: 1 +ici_context_ulysses_parallelism: 1 # Ulysses (head-exchange) dimension of the context parallelism; used by context_parallel_strategy='usp'. ici_context_autoregressive_parallelism: 1 ici_tensor_parallelism: 1 ici_tensor_sequence_parallelism: 1 @@ -1154,7 +1159,7 @@ cost_estimate_flops_bwd: -1 # -1 means using splash default cost estmiation, any dq_reduction_steps: 0 #the number of reduction steps. For now, only 3 or all the kv steps are supported. ### Determine if we want to use load balance for context parallelism context_parallel_load_balance: true -context_parallel_strategy: "all_gather" # "all_gather", "ring", or "ulysses" +context_parallel_strategy: "all_gather" # "all_gather", "ring", "ulysses", or "usp" context_parallel_reorder_strategy: "auto" # "auto", "dual_chunk_swap", or "striped" diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index acaf519a3d..704db003e6 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -1105,7 +1105,7 @@ class HardwareAndMesh(BaseModel): context_parallel_load_balance: bool = Field(True, description="Whether to use load balancing for context parallelism.") context_parallel_strategy: str = Field( "all_gather", - description="Strategy for context parallelism ('all_gather', 'ring', or 'ulysses').", + description="Strategy for context parallelism ('all_gather', 'ring', 'ulysses', or 'usp').", ) context_parallel_reorder_strategy: ReorderStrategy = Field( ReorderStrategy.AUTO, @@ -1159,6 +1159,13 @@ class LayoutAndSharding(BaseModel): ) data_sharding: Any = Field([], description="Sharding for input data.") context_sharding: str = Field("context", description="Physical axis name for context parallelism.") + ulysses_context_sharding: str = Field( + "context_ulysses", + description=( + "Physical axis name for the Ulysses head exchange under context_parallel_strategy='usp'. " + "context_sharding names the ring dimension; this names the all-to-all dimension." + ), + ) input_data_sharding_logical_axes: list[str] = Field( ["activation_embed_and_logits_batch", "activation_norm_length"], description="Logical axes for sharding input data.", @@ -1196,6 +1203,9 @@ class DcnParallelism(BaseModel): dcn_fsdp_transpose_parallelism: int = Field(1, description="DCN axis for FSDP transpose.") dcn_sequence_parallelism: int = Field(1, description="DCN axis for sequence parallelism (not recommended).") dcn_context_parallelism: int = Field(1, description="DCN axis for context parallelism.") + dcn_context_ulysses_parallelism: int = Field( + 1, description="DCN axis for the Ulysses dimension of USP context parallelism." + ) dcn_context_autoregressive_parallelism: int = Field(1, description="DCN axis for context autoregressive parallelism.") dcn_tensor_parallelism: int = Field(1, description="DCN axis for tensor parallelism (not recommended).") dcn_tensor_sequence_parallelism: int = Field( @@ -1215,6 +1225,9 @@ class IciParallelism(BaseModel): ici_fsdp_transpose_parallelism: int = Field(1, description="ICI axis for FSDP transpose.") ici_sequence_parallelism: int = Field(1, description="ICI axis for sequence parallelism.") ici_context_parallelism: int = Field(1, description="ICI axis for context parallelism.") + ici_context_ulysses_parallelism: int = Field( + 1, description="ICI axis for the Ulysses dimension of USP context parallelism." + ) ici_context_autoregressive_parallelism: int = Field(1, description="ICI axis for context autoregressive parallelism.") ici_tensor_parallelism: int = Field(1, description="ICI axis for tensor parallelism.") ici_tensor_sequence_parallelism: int = Field(1, description="ICI axis for tensor sequence parallelism.") @@ -2861,6 +2874,102 @@ def _validate_use_te_comm_gemm_overlap(self): "TE Collective GEMM operations are only supported for TE quantization recipes (i.e. starting with 'te_')." ) + def _validate_usp_context_parallelism(self): + """Validates the USP (Ulysses over ring) context parallelism configuration.""" + if self.context_parallel_strategy != "usp": + if self.ici_context_ulysses_parallelism != 1 or self.dcn_context_ulysses_parallelism != 1: + raise ValueError( + "ici/dcn_context_ulysses_parallelism was specified, but is only supported when " + "context_parallel_strategy='usp'." + ) + return + if self.hardware != "tpu": + raise ValueError("USP context parallelism (context_parallel_strategy='usp') is only supported on TPU.") + if self.context_sharding != "context": + raise ValueError("TPU USP attention requires context_sharding='context'.") + usp_sequence_axes = (self.context_sharding, self.ulysses_context_sharding) + for usp_axis in usp_sequence_axes: + if usp_axis not in self.mesh_axes: + raise ValueError(f"TPU USP attention requires mesh axis '{usp_axis}' in mesh_axes.") + if infer_cp_axes(self.logical_axis_rules) != usp_sequence_axes: + raise ValueError( + f"TPU USP attention requires activation_length to map to {usp_sequence_axes} in logical_axis_rules." + ) + if infer_cp_axes(self.logical_axis_rules_for_eval) != usp_sequence_axes: + raise ValueError( + f"TPU USP attention requires activation_length to map to {usp_sequence_axes} in logical_axis_rules_for_eval." + ) + usp_ring_size = self.ici_context_parallelism + usp_ulysses_size = self.ici_context_ulysses_parallelism + if ( + usp_ring_size <= 0 + or usp_ulysses_size <= 0 + or self.dcn_context_parallelism <= 0 + or self.dcn_context_ulysses_parallelism <= 0 + ): + raise ValueError( + "TPU USP attention requires explicit positive ici/dcn context parallelism values; " + "inferred (-1) sizes are not supported." + ) + if usp_ring_size <= 1: + raise ValueError("TPU USP attention requires ici_context_parallelism > 1 for the ring dimension.") + if usp_ulysses_size <= 1: + raise ValueError("TPU USP attention requires ici_context_ulysses_parallelism > 1 for the Ulysses dimension.") + if self.dcn_context_parallelism != 1 or self.dcn_context_ulysses_parallelism != 1: + raise ValueError("TPU USP attention does not support dcn context parallelism yet.") + if self.attention != "flash": + raise ValueError("TPU USP attention requires attention=flash.") + if not self.use_tokamax_splash: + raise ValueError("TPU USP attention requires use_tokamax_splash=True.") + if self.use_jax_splash: + raise ValueError("TPU USP attention requires use_jax_splash=False.") + if self.attention_type != "global": + raise ValueError("TPU USP attention is initially supported only for global causal attention.") + if self.packing: + raise ValueError("TPU USP attention does not support packing yet.") + if self.context_parallel_load_balance: + raise ValueError("TPU USP attention does not support context_parallel_load_balance=True.") + if self.use_ragged_attention: + raise ValueError("TPU USP attention does not support ragged attention.") + if self.attention_sink: + raise ValueError("TPU USP attention does not support attention sinks.") + if self.use_indexer: + raise ValueError("TPU USP attention does not support sparse indexer masks.") + if self.use_chunked_prefill: + raise ValueError("TPU USP attention does not support chunked prefill yet.") + if self.use_multimodal: + raise ValueError("TPU USP attention does not support multimodal attention.") + if self.enable_dropout and self.dropout_rate > 0.0: + raise ValueError("TPU USP attention does not support dropout yet.") + if self.dq_reduction_steps not in (0, 3): + raise ValueError("TPU USP attention requires dq_reduction_steps to be 0 or 3.") + if self.use_qk_clip: + raise ValueError("TPU USP attention does not support QK-Clip statistics yet.") + if self.mtp_num_layers > 0: + raise ValueError("TPU USP attention does not support multi-token prediction (mtp_num_layers > 0) yet.") + if self.sa_bwd_dkv_megacore: + raise ValueError("TPU USP attention does not support sa_bwd_dkv_megacore yet.") + if self.max_target_length % (usp_ring_size * usp_ulysses_size) != 0: + raise ValueError( + "TPU USP attention requires max_target_length " + f"({self.max_target_length}) to be divisible by the total context parallelism " + f"({usp_ring_size * usp_ulysses_size})." + ) + if self.max_target_length % (usp_ring_size * usp_ring_size) != 0: + raise ValueError("TPU USP attention requires max_target_length to be divisible by ici_context_parallelism squared.") + if self.num_query_heads % usp_ulysses_size != 0: + raise ValueError( + "TPU USP attention requires num_query_heads " + f"({self.num_query_heads}) to be divisible by ici_context_ulysses_parallelism ({usp_ulysses_size})." + ) + if self.num_kv_heads == 1: + raise ValueError("TPU USP attention does not support MQA with ici_context_ulysses_parallelism > 1.") + if self.num_kv_heads % usp_ulysses_size != 0: + raise ValueError( + "TPU USP attention requires num_kv_heads " + f"({self.num_kv_heads}) to be divisible by ici_context_ulysses_parallelism ({usp_ulysses_size})." + ) + def validate_num_moe_emb_chunks(self): """ Validates that num_moe_emb_chunks is used with supported settings. @@ -3644,6 +3753,8 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de ) if self.context_sharding not in ("context", "expert"): raise ValueError(f"Assigned context_sharding f{self.context_sharding} is not supported.") + if self.ulysses_context_sharding != "context_ulysses": + raise ValueError(f"Assigned ulysses_context_sharding {self.ulysses_context_sharding} is not supported.") if ( self.per_device_batch_size > 0 and (self.per_device_batch_size * self.max_target_length) % self.num_vocab_tiling != 0 @@ -3653,8 +3764,8 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de self, f"dcn_{self.context_sharding}_parallelism", 1 ) context_parallel_strategy = self.context_parallel_strategy.lower() - if context_parallel_strategy not in ("all_gather", "ring", "ulysses"): - raise ValueError("context_parallel_strategy must be one of 'all_gather', 'ring', or 'ulysses'.") + if context_parallel_strategy not in ("all_gather", "ring", "ulysses", "usp"): + raise ValueError("context_parallel_strategy must be one of 'all_gather', 'ring', 'ulysses', or 'usp'.") self.context_parallel_strategy = context_parallel_strategy if ( context_parallel_strategy == "ring" @@ -3711,10 +3822,10 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de raise ValueError("TPU Tokamax ring attention does not support QK-Clip statistics yet.") if self.enable_dropout and self.dropout_rate > 0.0: raise ValueError("TPU Tokamax ring attention does not support dropout yet.") - if context_parallel_strategy != "ring" and self.ring_scan_unroll != 1: + if context_parallel_strategy not in ("ring", "usp") and self.ring_scan_unroll != 1: raise ValueError( f"ring_scan_unroll={self.ring_scan_unroll} was specified, but is only supported when " - "context_parallel_strategy='ring'." + "context_parallel_strategy='ring' or 'usp'." ) if context_parallel_strategy == "ulysses": if self.hardware != "tpu": @@ -3779,6 +3890,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de "TPU Ulysses attention requires num_kv_heads " f"({self.num_kv_heads}) to be divisible by context_parallel_size ({context_parallel_size})." ) + self._validate_usp_context_parallelism() # STRIPED reorder strategy is a Transformer Engine feature and is GPU-only. # AUTO is resolved in training because test code paths may load the same # config but use a different reorder path. @@ -3801,6 +3913,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de * self.dcn_fsdp_transpose_parallelism * self.dcn_sequence_parallelism * self.dcn_context_parallelism + * self.dcn_context_ulysses_parallelism * self.dcn_tensor_parallelism * self.dcn_tensor_sequence_parallelism * self.dcn_expert_parallelism @@ -3975,6 +4088,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de "fsdp_transpose": self.ici_fsdp_transpose_parallelism, "sequence": self.ici_sequence_parallelism, "context": self.ici_context_parallelism, + "context_ulysses": self.ici_context_ulysses_parallelism, "context_autoregressive": self.ici_context_autoregressive_parallelism, "tensor": self.ici_tensor_parallelism, "tensor_sequence": self.ici_tensor_sequence_parallelism, @@ -3994,6 +4108,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de "fsdp_transpose": self.dcn_fsdp_transpose_parallelism, "sequence": self.dcn_sequence_parallelism, "context": self.dcn_context_parallelism, + "context_ulysses": self.dcn_context_ulysses_parallelism, "context_autoregressive": self.dcn_context_autoregressive_parallelism, "tensor": self.dcn_tensor_parallelism, "tensor_sequence": self.dcn_tensor_sequence_parallelism, diff --git a/src/maxtext/kernels/attention/ulysses_attention.py b/src/maxtext/kernels/attention/ulysses_attention.py index 79e384dea7..bcf7221b52 100644 --- a/src/maxtext/kernels/attention/ulysses_attention.py +++ b/src/maxtext/kernels/attention/ulysses_attention.py @@ -135,13 +135,14 @@ def validate_dkv_sharding( axis_names_kv: Any, dkv_dim_q: int, dkv_dim_kv: int, + attention_label: str, ) -> None: - """Validates that the head-dim/D_KV dimension stays local for Ulysses attention.""" + """Validates that the head-dim/D_KV dimension stays local for the head exchange.""" q_dkv_axes = sharding.mesh_axes_for_dim(axis_names_q[dkv_dim_q]) kv_dkv_axes = sharding.mesh_axes_for_dim(axis_names_kv[dkv_dim_kv]) if q_dkv_axes or kv_dkv_axes: raise ValueError( - "TPU Ulysses attention does not support sharding the D_KV/head-dim " + f"{attention_label} does not support sharding the D_KV/head-dim " f"dimension; got Q axes {q_dkv_axes} and K/V axes {kv_dkv_axes}." ) @@ -156,28 +157,29 @@ def validate_head_sharding( head_dim_q: int, head_dim_kv: int, ulysses_size: int, + attention_label: str, ) -> None: """Validates local head counts before the Ulysses head/sequence exchange.""" q_head_axes = sharding.mesh_axes_for_dim(axis_names_q[head_dim_q]) kv_head_axes = sharding.mesh_axes_for_dim(axis_names_kv[head_dim_kv]) - q_head_shards = sharding.mesh_axes_size(mesh, q_head_axes, label="TPU Ulysses attention") - kv_head_shards = sharding.mesh_axes_size(mesh, kv_head_axes, label="TPU Ulysses attention") + q_head_shards = sharding.mesh_axes_size(mesh, q_head_axes, label=attention_label) + kv_head_shards = sharding.mesh_axes_size(mesh, kv_head_axes, label=attention_label) if num_query_heads % q_head_shards != 0: raise ValueError( - "TPU Ulysses attention requires num_query_heads " + f"{attention_label} requires num_query_heads " f"({num_query_heads}) to be divisible by Q head shards ({q_head_shards})." ) if num_kv_heads % kv_head_shards != 0: raise ValueError( - "TPU Ulysses attention requires num_kv_heads " + f"{attention_label} requires num_kv_heads " f"({num_kv_heads}) to be divisible by KV head shards ({kv_head_shards})." ) if num_kv_heads == 1: - raise ValueError("TPU Ulysses attention does not support MQA with context_parallel_size > 1.") + raise ValueError(f"{attention_label} does not support MQA with a Ulysses exchange size > 1.") if q_head_axes != kv_head_axes: raise ValueError( - "TPU Ulysses attention requires Q and KV head sharding to match for MHA/GQA, " + f"{attention_label} requires Q and KV head sharding to match for MHA/GQA, " f"got Q head axes {q_head_axes} and KV head axes {kv_head_axes}." ) @@ -185,18 +187,18 @@ def validate_head_sharding( local_kv_heads = num_kv_heads // kv_head_shards if local_query_heads % local_kv_heads != 0: raise ValueError( - "TPU Ulysses attention requires local query heads " + f"{attention_label} requires local query heads " f"({local_query_heads}) to be divisible by local KV heads ({local_kv_heads})." ) if local_query_heads % ulysses_size != 0: raise ValueError( - "TPU Ulysses attention requires local query heads " - f"({local_query_heads}) to be divisible by context_parallel_size ({ulysses_size})." + f"{attention_label} requires local query heads " + f"({local_query_heads}) to be divisible by the Ulysses exchange size ({ulysses_size})." ) if local_kv_heads % ulysses_size != 0: raise ValueError( - "TPU Ulysses attention requires local KV heads " - f"({local_kv_heads}) to be divisible by context_parallel_size ({ulysses_size})." + f"{attention_label} requires local KV heads " + f"({local_kv_heads}) to be divisible by the Ulysses exchange size ({ulysses_size})." ) diff --git a/src/maxtext/kernels/attention/usp_attention.py b/src/maxtext/kernels/attention/usp_attention.py new file mode 100644 index 0000000000..ebb6c909ed --- /dev/null +++ b/src/maxtext/kernels/attention/usp_attention.py @@ -0,0 +1,167 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""USP (Ulysses-over-ring) attention layout helpers.""" + +from __future__ import annotations + +from typing import Any + +import jax + +from maxtext.common.common_types import MODEL_MODE_TRAIN +from maxtext.kernels.attention import tokamax_ring_attention +from maxtext.kernels.attention import ulysses_attention +from maxtext.utils import sharding + + +def is_context_parallel_usp_requested(config: Any) -> bool: + """Returns True when the config requests USP context parallelism.""" + return config.context_parallel_strategy == "usp" + + +def validate_usp_runtime( + *, + model_mode: str, + use_ragged_attention: bool = False, + previous_chunk: Any = None, + sinks: Any = None, + indexer_mask: Any = None, + bidirectional_mask: Any = None, + record_max_logits: bool = False, +) -> None: + """Validates runtime-only constraints for the USP path.""" + if model_mode != MODEL_MODE_TRAIN: + raise ValueError("TPU USP attention is supported only for train mode.") + if use_ragged_attention: + raise ValueError("TPU USP attention does not support ragged attention.") + if previous_chunk is not None: + raise ValueError("TPU USP attention does not support chunked prefill yet.") + if sinks is not None: + raise ValueError("TPU USP attention does not support attention sinks.") + if indexer_mask is not None: + raise ValueError("TPU USP attention does not support indexer masks.") + if bidirectional_mask is not None: + raise ValueError("TPU USP attention does not support bidirectional masks.") + if record_max_logits: + raise NotImplementedError("TPU USP attention does not support record_max_logits yet.") + + +def call_usp_attention( + query: Any, + key: Any, + value: Any, + decoder_segment_ids_q: Any, + ring_kernel: Any, + ulysses_axis: str, +): + """Runs ring attention over the Ulysses-exchanged operands and restores the layout.""" + query = ulysses_attention.ulysses_all_to_all(query, ulysses_axis) + key = ulysses_attention.ulysses_all_to_all(key, ulysses_axis) + value = ulysses_attention.ulysses_all_to_all(value, ulysses_axis) + if decoder_segment_ids_q is not None: + # One gather serves both ring operands; the result stays sequence-sharded + # over the ring axis, the layout the ring kernel expects. + ring_segment_ids = jax.lax.all_gather(decoder_segment_ids_q, ulysses_axis, axis=1, tiled=True) + else: + ring_segment_ids = None + attention_output = tokamax_ring_attention.call_ring_attention( + query, + key, + value, + ring_segment_ids, + ring_segment_ids, + ring_kernel, + ) + return ulysses_attention.inverse_ulysses_all_to_all(attention_output, ulysses_axis) + + +def with_usp_sequence_axes(axis_names: Any, ring_axis: str, ulysses_axis: str, sequence_dim: int) -> Any: + """Returns axis names with the sequence dimension set to the ring and Ulysses axes.""" + if axis_names is None: + return None + if len(axis_names) <= sequence_dim: + raise ValueError("TPU USP attention expects a sequence sharding dimension.") + expected_axes = (ring_axis, ulysses_axis) + existing_sequence_axes = sharding.mesh_axes_for_dim(axis_names[sequence_dim]) + if existing_sequence_axes and existing_sequence_axes != expected_axes: + raise ValueError( + "TPU USP attention expects the existing sequence sharding to be " + f"unsharded or exactly {expected_axes}, got {existing_sequence_axes}." + ) + return sharding.with_axis_on_dim(axis_names, expected_axes, sequence_dim) + + +def _validate_usp_axes_only_on_sequence( + axis_names: Any, + *, + tensor_name: str, + sequence_dim: int, + ring_axis: str, + ulysses_axis: str, +) -> None: + """Raises if a USP mesh axis appears outside the sequence dimension.""" + for dim, axis_name in enumerate(axis_names): + if dim == sequence_dim: + continue + dim_axes = sharding.mesh_axes_for_dim(axis_name) + for axis in (ring_axis, ulysses_axis): + if axis in dim_axes: + raise ValueError( + "TPU USP attention requires the context axes to appear only " + f"on the sequence dimension; got {axis!r} on {tensor_name} dim {dim}." + ) + + +def validate_usp_mesh_axes( + *, + axis_names_q: Any, + axis_names_kv: Any, + sequence_dim_q: int, + sequence_dim_kv: int, + mesh: Any, + ring_axis: str, + ulysses_axis: str, +) -> None: + """Validates sequence sharding before the USP exchange.""" + if ring_axis == ulysses_axis: + raise ValueError("TPU USP attention requires context_sharding and ulysses_context_sharding to differ.") + for axis in (ring_axis, ulysses_axis): + if axis not in mesh.shape: + raise ValueError(f"TPU USP attention requires mesh axis {axis!r} to exist.") + _validate_usp_axes_only_on_sequence( + axis_names_q, + tensor_name="Q", + sequence_dim=sequence_dim_q, + ring_axis=ring_axis, + ulysses_axis=ulysses_axis, + ) + _validate_usp_axes_only_on_sequence( + axis_names_kv, + tensor_name="K/V", + sequence_dim=sequence_dim_kv, + ring_axis=ring_axis, + ulysses_axis=ulysses_axis, + ) + + expected_axes = (ring_axis, ulysses_axis) + q_sequence_axes = sharding.mesh_axes_for_dim(axis_names_q[sequence_dim_q]) + kv_sequence_axes = sharding.mesh_axes_for_dim(axis_names_kv[sequence_dim_kv]) + if q_sequence_axes != expected_axes: + raise ValueError( + f"TPU USP attention requires Q sequence sharding to be exactly {expected_axes}, got {q_sequence_axes}." + ) + if kv_sequence_axes != expected_axes: + raise ValueError( + f"TPU USP attention requires K/V sequence sharding to be exactly {expected_axes}, got {kv_sequence_axes}." + ) diff --git a/src/maxtext/layers/attention_op.py b/src/maxtext/layers/attention_op.py index 25dc782ab7..c92b23fd9e 100644 --- a/src/maxtext/layers/attention_op.py +++ b/src/maxtext/layers/attention_op.py @@ -66,6 +66,7 @@ from maxtext.kernels.attention import jax_flash_attention from maxtext.kernels.attention import tokamax_ring_attention from maxtext.kernels.attention import ulysses_attention +from maxtext.kernels.attention import usp_attention from maxtext.kernels.attention.ragged_attention import ragged_gqa from maxtext.kernels.attention.ragged_attention import ragged_mha from maxtext.layers import nnx_wrappers @@ -702,12 +703,70 @@ def __init__( head_dim_q=1, head_dim_kv=1, ulysses_size=self.mesh.shape[context_axis], + attention_label="TPU Ulysses attention", ) ulysses_attention.validate_dkv_sharding( axis_names_q=axis_names_q, axis_names_kv=axis_names_kv, dkv_dim_q=3, dkv_dim_kv=3, + attention_label="TPU Ulysses attention", + ) + if self.attention_kernel == "flash" and usp_attention.is_context_parallel_usp_requested(self.config): + target_hardware = self.mesh.devices[(0,) * self.mesh.devices.ndim].platform + if target_hardware != "tpu": + raise ValueError("USP context parallelism (context_parallel_strategy='usp') is only supported on TPU.") + if not self.config.use_tokamax_splash: + raise ValueError("TPU USP attention requires use_tokamax_splash=True.") + if self.config.use_jax_splash: + raise ValueError("TPU USP attention requires use_jax_splash=False.") + if self.attention_type != AttentionType.GLOBAL: + raise ValueError("TPU USP attention is initially supported only for global causal attention.") + if self.config.enable_dropout and self.dropout_rate > 0.0: + raise ValueError("TPU USP attention does not support dropout yet.") + if self.use_ragged_attention: + raise ValueError("TPU USP attention does not support ragged attention.") + + ring_axis = self.config.context_sharding + ulysses_axis = self.config.ulysses_context_sharding + axis_names_q = self._logical_to_mesh_axes(self.flash_axis_names_q) + axis_names_kv = self._logical_to_mesh_axes(self.flash_axis_names_kv) + axis_names_kv = usp_attention.with_usp_sequence_axes( + axis_names_kv, + ring_axis, + ulysses_axis, + sequence_dim=2, + ) + usp_attention.validate_usp_mesh_axes( + axis_names_q=axis_names_q, + axis_names_kv=axis_names_kv, + sequence_dim_q=2, + sequence_dim_kv=2, + mesh=self.mesh, + ring_axis=ring_axis, + ulysses_axis=ulysses_axis, + ) + if self.mesh.shape[ring_axis] <= 1: + raise ValueError("TPU USP attention requires a context parallel mesh axis larger than one.") + if self.mesh.shape[ulysses_axis] <= 1: + raise ValueError("TPU USP attention requires a context ulysses parallel mesh axis larger than one.") + ulysses_attention.validate_head_sharding( + axis_names_q=axis_names_q, + axis_names_kv=axis_names_kv, + mesh=self.mesh, + num_query_heads=self.num_query_heads, + num_kv_heads=self.num_kv_heads, + head_dim_q=1, + head_dim_kv=1, + ulysses_size=self.mesh.shape[ulysses_axis], + attention_label="TPU USP attention", + ) + ulysses_attention.validate_dkv_sharding( + axis_names_q=axis_names_q, + axis_names_kv=axis_names_kv, + dkv_dim_q=3, + dkv_dim_kv=3, + attention_label="TPU USP attention", ) def maybe_create_nnx(einsum, *args): @@ -1243,6 +1302,28 @@ def _validate_tpu_ulysses_runtime( record_max_logits=record_max_logits, ) + def _validate_tpu_usp_runtime( + self, + *, + model_mode: str, + previous_chunk: Any = None, + bidirectional_mask: Any = None, + sinks: Array | None = None, + indexer_mask: Array | None = None, + use_ragged_attention: bool = False, + record_max_logits: bool = False, + ) -> None: + """Validates runtime constraints for the TPU USP path.""" + usp_attention.validate_usp_runtime( + model_mode=model_mode, + previous_chunk=previous_chunk, + sinks=sinks, + indexer_mask=indexer_mask, + use_ragged_attention=use_ragged_attention, + bidirectional_mask=bidirectional_mask, + record_max_logits=record_max_logits, + ) + def apply_attention( self, query: Array, @@ -1278,6 +1359,11 @@ def apply_attention( raise ValueError("Ulysses context parallelism (context_parallel_strategy='ulysses') is only supported on TPU.") if self.attention_kernel != "flash": raise ValueError("TPU Ulysses attention requires attention_kernel='flash'.") + if usp_attention.is_context_parallel_usp_requested(self.config): + if target_hardware != "tpu": + raise ValueError("USP context parallelism (context_parallel_strategy='usp') is only supported on TPU.") + if self.attention_kernel != "flash": + raise ValueError("TPU USP attention requires attention_kernel='flash'.") if use_ragged_attention and model_mode == MODEL_MODE_AUTOREGRESSIVE: if lengths is None: @@ -1524,6 +1610,7 @@ def tpu_flash_attention( use_tokamax_ring = tokamax_ring_attention.is_context_parallel_ring_requested(self.config) use_ulysses = ulysses_attention.is_context_parallel_ulysses_requested(self.config) + use_usp = usp_attention.is_context_parallel_usp_requested(self.config) cp_size = self.mesh.shape.get(self.config.context_sharding, 1) load_balanced_context_parallel = self.config.context_parallel_load_balance if use_tokamax_ring: @@ -1546,6 +1633,16 @@ def tpu_flash_attention( use_ragged_attention=use_ragged_attention, record_max_logits=record_max_logits, ) + elif use_usp: + self._validate_tpu_usp_runtime( + model_mode=model_mode, + previous_chunk=previous_chunk, + bidirectional_mask=bidirectional_mask, + sinks=sinks, + indexer_mask=indexer_mask, + use_ragged_attention=use_ragged_attention, + record_max_logits=record_max_logits, + ) # Transpose to ('batch', 'heads', 'length', 'kv') query = jnp.transpose(query, axes=(0, 2, 1, 3)) @@ -1596,6 +1693,27 @@ def tpu_flash_attention( context_axis, sequence_dim=1, ) + elif use_usp: + ring_axis = self.config.context_sharding + ulysses_axis = self.config.ulysses_context_sharding + segment_axis_names_q = usp_attention.with_usp_sequence_axes( + segment_axis_names_q, + ring_axis, + ulysses_axis, + sequence_dim=1, + ) + axis_names_kv = usp_attention.with_usp_sequence_axes( + axis_names_kv, + ring_axis, + ulysses_axis, + sequence_dim=2, + ) + segment_axis_names_kv = usp_attention.with_usp_sequence_axes( + segment_axis_names_kv, + ring_axis, + ulysses_axis, + sequence_dim=1, + ) devices_in_data_fsdp = self.mesh.shape.get("data", 1) * self.mesh.shape.get("fsdp", 1) assert (query.shape[0] / devices_in_data_fsdp).is_integer(), ( @@ -1656,7 +1774,7 @@ def create_sa_config(config, query, key, attn_logits_soft_cap): ) return sa_config - if use_tokamax_ring: + if use_tokamax_ring or use_usp: sa_config, splash_kernel, segment_axis_names_splash_kernel = ( tokamax_ring_attention.make_sharded_ring_attention_kernel( self.config, @@ -1754,7 +1872,7 @@ def wrap_ulysses_splash_kernel(single_head_mask): ) max_logit_value = None - if not use_tokamax_ring and not use_ulysses and self.config.use_tokamax_splash: + if not use_tokamax_ring and not use_ulysses and not use_usp and self.config.use_tokamax_splash: # Create mask single_head_mask = mask # tokamax now just uses a single mask and assumes broadcast to all heads if self.config.use_max_logit_estimate > 0: @@ -1778,11 +1896,11 @@ def wrap_tokamax_splash_kernel(single_head_mask): splash_kernel = wrap_tokamax_splash_kernel(single_head_mask) segment_axis_names_splash_kernel = self._logical_to_mesh_axes((Q_LENGTH,)) splash_kernel = self._maybe_shard_with_pspec(splash_kernel, segment_axis_names_splash_kernel) - elif not use_tokamax_ring and not use_ulysses and self.config.use_jax_splash: + elif not use_tokamax_ring and not use_ulysses and not use_usp and self.config.use_jax_splash: if self.config.use_max_logit_estimate > 0: sa_config = dataclasses.replace(sa_config, max_logit_const=self.config.use_max_logit_estimate) segment_axis_names_splash_kernel = nn.logical_to_mesh_axes((Q_LENGTH,)) - elif not use_tokamax_ring and not use_ulysses: + elif not use_tokamax_ring and not use_ulysses and not use_usp: # Create multi-head mask multi_head_mask = splash_attention_mask.MultiHeadMask(masks=(mask,) * query.shape[1]) @@ -1823,7 +1941,9 @@ def wrap_jax_splash_kernel(multi_head_mask, shard_head_size=1): # sequence-sharded and K/V are replicated. For the Tokamax ring path Q, K, # V, and segment IDs are all sequence-sharded over the context axis. # For Ulysses Q/K/V are sequence-sharded at the boundary and head-sharded - # inside the local Splash call. + # inside the local Splash call. USP performs the same Ulysses exchange over + # its ulysses axis and runs the ring kernel over the ring axis within each + # head subset. if record_max_logits: # max_logits will share similar sharding as query but last dim is unrelated to model @@ -1908,6 +2028,17 @@ def wrap_flash_attention( attention_output = ulysses_attention.inverse_ulysses_all_to_all(attention_output, context_axis) return attention_output, None + if use_usp: + attention_output = usp_attention.call_usp_attention( + query, + key, + value, + decoder_segment_ids_q, + splash_kernel, + ulysses_axis, + ) + return attention_output, None + # The load-balanced all-gather path restores K/V to contiguous order # before calling Splash attention. if cp_size > 1 and load_balanced_context_parallel: diff --git a/src/maxtext/utils/maxtext_utils.py b/src/maxtext/utils/maxtext_utils.py index f162fc8275..9e7311a70d 100644 --- a/src/maxtext/utils/maxtext_utils.py +++ b/src/maxtext/utils/maxtext_utils.py @@ -2169,6 +2169,7 @@ def create_device_mesh(config, devices=None): "fsdp_transpose": getattr(config, "ici_fsdp_transpose_parallelism", 1), "sequence": getattr(config, "ici_sequence_parallelism", 1), "context": getattr(config, "ici_context_parallelism", 1), + "context_ulysses": getattr(config, "ici_context_ulysses_parallelism", 1), "context_autoregressive": getattr(config, "ici_context_autoregressive_parallelism", 1), "tensor": getattr(config, "ici_tensor_parallelism", 1), "tensor_sequence": getattr(config, "ici_tensor_sequence_parallelism", 1), @@ -2196,6 +2197,7 @@ def create_device_mesh(config, devices=None): "fsdp_transpose": getattr(config, "dcn_fsdp_transpose_parallelism", 1), "sequence": getattr(config, "dcn_sequence_parallelism", 1), "context": getattr(config, "dcn_context_parallelism", 1), + "context_ulysses": getattr(config, "dcn_context_ulysses_parallelism", 1), "context_autoregressive": getattr(config, "dcn_context_autoregressive_parallelism", 1), "tensor": getattr(config, "dcn_tensor_parallelism", 1), "tensor_sequence": getattr(config, "dcn_tensor_sequence_parallelism", 1), diff --git a/src/maxtext/utils/sharding.py b/src/maxtext/utils/sharding.py index 3ebaa21610..914c938d27 100644 --- a/src/maxtext/utils/sharding.py +++ b/src/maxtext/utils/sharding.py @@ -452,6 +452,7 @@ def _get_nontrival_mesh_axes(mesh): "fsdp_transpose", "sequence", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", diff --git a/tests/unit/attention_test.py b/tests/unit/attention_test.py index dc7ea03f43..8a4565f027 100644 --- a/tests/unit/attention_test.py +++ b/tests/unit/attention_test.py @@ -1942,6 +1942,166 @@ def attention_loss(x, pos, seg): ) self.assertLen(hlo_test_utils.collective_lines(hlo_text, "collective-permute"), 0) + def _usp_test_config(self): + return pyconfig.initialize( + [sys.argv[0], get_test_config_path()], + **self.config_arguments, + attention="flash", + context_parallel_strategy="usp", + context_parallel_load_balance=False, + ici_context_parallelism=2, + ici_context_ulysses_parallelism=2, + use_tokamax_splash=True, + use_jax_splash=False, + packing=False, + dtype="float32", + ) + + @pytest.mark.tpu_only + def test_tpu_flash_attention_usp_context_parallel(self): + """Test equivalence between dot_product and flash attention + USP context parallelism""" + + cfg_cp = self._usp_test_config() + devices_array_cp = maxtext_utils.create_device_mesh(cfg_cp) + mesh_cp = Mesh(devices_array_cp, cfg_cp.mesh_axes) + lnx, decoder_segment_ids, decoder_positions = self.get_data(cfg_cp.dtype) + attention_as_mha_generic, attention_as_mha_flash_cp = self._ulysses_test_modules(cfg_cp, mesh_cp, lnx) + mha_generic_output, _ = attention_as_mha_generic( + lnx, + lnx, + decoder_segment_ids=decoder_segment_ids, + inputs_positions=decoder_positions, + deterministic=True, + model_mode=MODEL_MODE_TRAIN, + ) + nnx.update(attention_as_mha_flash_cp, nnx.state(attention_as_mha_generic)) + + mha_generic_flash_cp_output = attention_test_util.forward_with_context_expert_parallelism( + cfg_cp, + mesh_cp, + attention_as_mha_flash_cp, + lnx, + decoder_segment_ids, + decoder_positions, + ) + + mha_generic_output = jax.device_get(mha_generic_output) + mha_generic_flash_cp_output = jax.device_get(mha_generic_flash_cp_output) + + self.assertTrue( + jax.numpy.allclose(mha_generic_output, mha_generic_flash_cp_output, rtol=1e-02, atol=1e-02, equal_nan=False), + msg="Logits from generic dot product and flash attention + USP context parallelism are not close.", + ) + + @pytest.mark.tpu_only + def test_tpu_flash_attention_usp_context_parallel_grad(self): + """Test input-gradient equivalence between dot_product and flash attention + USP context parallelism""" + + cfg_cp = self._usp_test_config() + devices_array_cp = maxtext_utils.create_device_mesh(cfg_cp) + mesh_cp = Mesh(devices_array_cp, cfg_cp.mesh_axes) + lnx, decoder_segment_ids, decoder_positions = self.get_data(cfg_cp.dtype) + attention_as_mha_generic, attention_as_mha_flash_cp = self._ulysses_test_modules(cfg_cp, mesh_cp, lnx) + nnx.update(attention_as_mha_flash_cp, nnx.state(attention_as_mha_generic)) + + def generic_loss(lnx): + output, _ = attention_as_mha_generic( + lnx, + lnx, + decoder_segment_ids=decoder_segment_ids, + inputs_positions=decoder_positions, + deterministic=True, + model_mode=MODEL_MODE_TRAIN, + ) + return jnp.mean(output.astype(jnp.float32) ** 2) + + def usp_loss(lnx): + output, _ = attention_as_mha_flash_cp( + lnx, + lnx, + decoder_segment_ids=decoder_segment_ids, + inputs_positions=decoder_positions, + deterministic=True, + model_mode=MODEL_MODE_TRAIN, + ) + return jnp.mean(output.astype(jnp.float32) ** 2) + + generic_grad = jax.grad(generic_loss)(lnx) + with jax.set_mesh(mesh_cp), nn_partitioning.axis_rules(cfg_cp.logical_axis_rules): + usp_grad = jax.grad(usp_loss)(lnx) + generic_grad = jax.device_get(generic_grad) + usp_grad = jax.device_get(usp_grad) + + self.assertTrue( + jax.numpy.allclose(generic_grad, usp_grad, rtol=1e-02, atol=1e-07, equal_nan=False), + msg="Input gradients from generic dot product and flash attention + USP context parallelism are not close.", + ) + + @pytest.mark.tpu_only + def test_tpu_flash_attention_usp_hlo_uses_all_to_all_and_permute(self): + """Checks compiled TPU USP attention HLO uses all-to-all and collective-permute.""" + + cfg_cp = self._usp_test_config() + devices_array_cp = maxtext_utils.create_device_mesh(cfg_cp) + mesh_cp = Mesh(devices_array_cp, cfg_cp.mesh_axes) + lnx, decoder_segment_ids, decoder_positions = self.get_data(cfg_cp.dtype) + _, attention_as_mha_flash_cp = self._ulysses_test_modules(cfg_cp, mesh_cp, lnx) + + def attention_forward(x, pos, seg): + output, _ = attention_as_mha_flash_cp( + x, + x, + decoder_segment_ids=seg, + inputs_positions=pos, + deterministic=True, + model_mode=MODEL_MODE_TRAIN, + ) + return output + + def attention_loss(x, pos, seg): + return jnp.sum(attention_forward(x, pos, seg).astype(jnp.float32)) + + hlo_texts = [] + for lowered_fn in (attention_forward, jax.grad(attention_loss)): + # The mesh and axis-rules contexts wrap the jit from outside because + # jax.set_mesh raises inside a traced function, and the output keeps its + # natural sequence sharding so the only full-sequence gathers in the + # program are the ones the attention path itself emits. + with jax.set_mesh(mesh_cp), nn_partitioning.axis_rules(cfg_cp.logical_axis_rules): + input_sharding = NamedSharding( + mesh_cp, + nn_partitioning.logical_to_mesh_axes( + ("activation_batch", "activation_length", "activation_embed"), nn_partitioning.get_axis_rules() + ), + ) + metadata_sharding = NamedSharding( + mesh_cp, nn_partitioning.logical_to_mesh_axes((None, "activation_length"), nn_partitioning.get_axis_rules()) + ) + lowered = jax.jit(lowered_fn).lower( + jax.device_put(lnx, input_sharding), + jax.device_put(decoder_positions, metadata_sharding), + jax.device_put(decoder_segment_ids, metadata_sharding), + ) + hlo_texts.append(lowered.compile().as_text()) + + ring_local_sequence_length = cfg_cp.max_target_length // cfg_cp.ici_context_parallelism + sequence_lengths = (cfg_cp.max_target_length, ring_local_sequence_length) + for hlo_text in hlo_texts: + self.assertGreater(len(hlo_test_utils.collective_lines(hlo_text, "all-to-all")), 0) + self.assertGreater(len(hlo_test_utils.collective_lines(hlo_text, "collective-permute")), 0) + self.assertLen(hlo_test_utils.attention_sequence_all_gather_lines(hlo_text, sequence_lengths), 0) + # The int32 segment-ID gather over the Ulysses axis spans one ring-local + # sequence block; it is the only intended sequence all-gather. + self.assertGreater( + len( + hlo_test_utils.attention_sequence_all_gather_lines(hlo_text, (ring_local_sequence_length,), dtypes=("s32",)) + ), + 0, + ) + self.assertLen( + hlo_test_utils.attention_sequence_all_gather_lines(hlo_text, (cfg_cp.max_target_length,), dtypes=("s32",)), 0 + ) + @pytest.mark.tpu_only def test_dot_product_cache_axis_order(self): all_axis_orders = tuple(itertools.permutations(range(4))) diff --git a/tests/unit/configs_value_test.py b/tests/unit/configs_value_test.py index 89db8cd7b0..bc7791ddb9 100644 --- a/tests/unit/configs_value_test.py +++ b/tests/unit/configs_value_test.py @@ -408,6 +408,128 @@ def test_tpu_ulysses_config_validation_rejects_unsupported_configs(self): with self.assertRaisesRegex((ValueError, pydantic.ValidationError), expected_regex): pyconfig.initialize(argv) + def test_tpu_usp_config_validation_accepts_initial_config(self): + argv = [ + "", + _BASE_CONFIG_PATH, + "run_name=test", + "attention=flash", + "use_tokamax_splash=True", + "use_jax_splash=False", + "context_parallel_strategy=usp", + "context_parallel_load_balance=False", + "ici_context_parallelism=2", + "ici_context_ulysses_parallelism=2", + "ring_scan_unroll=2", + "hardware=tpu", + "packing=False", + "dataset_type=synthetic", + "skip_jax_distributed_system=True", + ] + mock_devices = [unittest.mock.MagicMock(slice_index=0) for _ in range(8)] + with unittest.mock.patch("jax.devices", return_value=mock_devices): + config = pyconfig.initialize(argv) + + self.assertEqual(config.context_parallel_strategy, "usp") + self.assertEqual(config.ici_context_parallelism, 2) + self.assertEqual(config.ici_context_ulysses_parallelism, 2) + self.assertEqual(config.ring_scan_unroll, 2) + self.assertEqual(config.ulysses_context_sharding, "context_ulysses") + context_ulysses_index = config.mesh_axes.index("context_ulysses") + self.assertEqual(context_ulysses_index, config.mesh_axes.index("context") + 1) + self.assertEqual(config.ici_parallelism[context_ulysses_index], 2) + self.assertEqual(types.infer_cp_axes(config.logical_axis_rules), ("context", "context_ulysses")) + self.assertEqual(types.infer_cp_axes(config.logical_axis_rules_for_eval), ("context", "context_ulysses")) + + def test_context_ulysses_parallelism_requires_usp(self): + argv = [ + "", + _BASE_CONFIG_PATH, + "run_name=test", + "ici_context_ulysses_parallelism=2", + "dataset_type=synthetic", + "skip_jax_distributed_system=True", + ] + mock_devices = [unittest.mock.MagicMock(slice_index=0) for _ in range(8)] + with unittest.mock.patch("jax.devices", return_value=mock_devices): + with self.assertRaisesRegex( + (ValueError, pydantic.ValidationError), "only supported when context_parallel_strategy='usp'" + ): + pyconfig.initialize(argv) + + def test_tpu_usp_config_validation_rejects_unsupported_configs(self): + base_args = [ + "", + _BASE_CONFIG_PATH, + "run_name=test", + "attention=flash", + "use_tokamax_splash=True", + "use_jax_splash=False", + "context_parallel_strategy=usp", + "context_parallel_load_balance=False", + "ici_context_parallelism=2", + "ici_context_ulysses_parallelism=2", + "hardware=tpu", + "packing=False", + "dataset_type=synthetic", + "skip_jax_distributed_system=True", + ] + cases = [ + (["context_parallel_load_balance=True"], ["context_parallel_load_balance=False"], "load_balance"), + (["packing=True", "dataset_type=tfds"], ["packing=False", "dataset_type=synthetic"], "packing"), + (["attention=dot_product"], ["attention=flash"], "attention=flash"), + (["use_tokamax_splash=False"], ["use_tokamax_splash=True"], "use_tokamax_splash"), + (["use_jax_splash=True"], ["use_jax_splash=False"], "use_jax_splash"), + (["attention_type=mla"], [], "global causal attention"), + (["use_ragged_attention=True"], [], "ragged attention"), + (["attention_sink=True"], [], "attention sinks"), + (["use_indexer=True", "q_lora_rank=1"], [], "sparse indexer"), + (["use_chunked_prefill=True"], [], "chunked prefill"), + (["use_multimodal=True"], [], "multimodal"), + (["dropout_rate=0.1"], [], "dropout"), + (["dq_reduction_steps=2"], [], "dq_reduction_steps"), + (["use_qk_clip=True"], [], "QK-Clip"), + (["context_sharding=expert"], [], "context_sharding"), + (["ulysses_context_sharding=expert"], [], "ulysses_context_sharding"), + (["custom_mesh_and_rule=pure-fsdp"], [], "mesh axis 'context' in"), + (["custom_mesh_and_rule=cp-as-ep"], [], "mesh axis 'context_ulysses' in"), + (["logical_axis_rules=[['activation_length',['context']]]"], [], r"in logical_axis_rules\."), + (["custom_mesh_and_rule_for_eval=pure-fsdp"], [], "logical_axis_rules_for_eval"), + (["ici_context_parallelism=1"], ["ici_context_parallelism=2"], "ring dimension"), + (["ici_context_ulysses_parallelism=1"], ["ici_context_ulysses_parallelism=2"], "Ulysses dimension"), + (["ici_context_parallelism=-1"], ["ici_context_parallelism=2"], "explicit positive"), + (["ici_context_ulysses_parallelism=-1"], ["ici_context_ulysses_parallelism=2"], "explicit positive"), + (["dcn_context_parallelism=2"], [], "dcn context parallelism"), + (["dcn_context_ulysses_parallelism=2"], [], "dcn context parallelism"), + (["dcn_context_parallelism=-1"], [], "explicit positive"), + (["dcn_context_ulysses_parallelism=-1"], [], "explicit positive"), + (["mtp_num_layers=1"], [], "multi-token prediction"), + (["sa_bwd_dkv_megacore=True"], [], "sa_bwd_dkv_megacore"), + (["max_target_length=2050"], [], "total context parallelism"), + (["ici_context_parallelism=4", "max_target_length=2056"], ["ici_context_parallelism=2"], "squared"), + ( + ["base_num_query_heads=18", "ici_context_ulysses_parallelism=4"], + ["ici_context_ulysses_parallelism=2"], + "requires num_query_heads", + ), + (["base_num_kv_heads=1"], [], "MQA"), + ( + ["base_num_kv_heads=10", "ici_context_ulysses_parallelism=4"], + ["ici_context_ulysses_parallelism=2"], + "requires num_kv_heads", + ), + (["hardware=gpu"], ["hardware=tpu"], "only supported on TPU"), + (["hardware=cpu"], ["hardware=tpu"], "only supported on TPU"), + ] + mock_devices = [unittest.mock.MagicMock(slice_index=0) for _ in range(8)] + for bad_args, args_to_remove, expected_regex in cases: + with self.subTest(bad_args=bad_args): + argv = [arg for arg in base_args if arg not in args_to_remove] + argv.extend(bad_args) + with unittest.mock.patch("jax.devices", return_value=mock_devices): + with self.assertRaisesRegex((ValueError, pydantic.ValidationError), expected_regex): + pyconfig.initialize(argv) + def test_load_balanced_chunk_context_parallel_config(self): argv = [ "", diff --git a/tests/unit/ulysses_attention_test.py b/tests/unit/ulysses_attention_test.py index be5995adbc..963cff31d1 100644 --- a/tests/unit/ulysses_attention_test.py +++ b/tests/unit/ulysses_attention_test.py @@ -126,6 +126,7 @@ def test_layout_validators_reject_invalid_shardings(self): axis_names_kv=(None, None, "context", None), dkv_dim_q=3, dkv_dim_kv=3, + attention_label="TPU Ulysses attention", ), ), ( @@ -139,6 +140,7 @@ def test_layout_validators_reject_invalid_shardings(self): head_dim_q=1, head_dim_kv=1, ulysses_size=4, + attention_label="TPU Ulysses attention", ), ), ( @@ -152,6 +154,7 @@ def test_layout_validators_reject_invalid_shardings(self): head_dim_q=1, head_dim_kv=1, ulysses_size=4, + attention_label="TPU Ulysses attention", ), ), ( @@ -165,10 +168,11 @@ def test_layout_validators_reject_invalid_shardings(self): head_dim_q=1, head_dim_kv=1, ulysses_size=4, + attention_label="TPU Ulysses attention", ), ), ( - r"local query heads \(8\) to be divisible by context_parallel_size", + r"local query heads \(8\) to be divisible by the Ulysses exchange size", lambda: ulysses_attention.validate_head_sharding( axis_names_q=(None, None, "context", None), axis_names_kv=(None, None, "context", None), @@ -178,6 +182,7 @@ def test_layout_validators_reject_invalid_shardings(self): head_dim_q=1, head_dim_kv=1, ulysses_size=16, + attention_label="TPU Ulysses attention", ), ), ] @@ -198,6 +203,7 @@ def test_validate_head_sharding_uses_local_heads_after_tensor_sharding(self): head_dim_q=1, head_dim_kv=1, ulysses_size=4, + attention_label="TPU Ulysses attention", ) with self.assertRaisesRegex(ValueError, "local KV heads"): @@ -210,6 +216,7 @@ def test_validate_head_sharding_uses_local_heads_after_tensor_sharding(self): head_dim_q=1, head_dim_kv=1, ulysses_size=4, + attention_label="TPU Ulysses attention", ) def test_validate_head_sharding_rejects_mqa(self): @@ -225,6 +232,7 @@ def test_validate_head_sharding_rejects_mqa(self): head_dim_q=1, head_dim_kv=1, ulysses_size=4, + attention_label="TPU Ulysses attention", ) def test_validate_head_sharding_requires_q_and_kv_head_axes_to_match(self): @@ -240,6 +248,7 @@ def test_validate_head_sharding_requires_q_and_kv_head_axes_to_match(self): head_dim_q=1, head_dim_kv=1, ulysses_size=4, + attention_label="TPU Ulysses attention", ) def test_ulysses_all_to_all_moves_heads_to_sequence(self): diff --git a/tests/unit/usp_attention_test.py b/tests/unit/usp_attention_test.py new file mode 100644 index 0000000000..6a1853dde9 --- /dev/null +++ b/tests/unit/usp_attention_test.py @@ -0,0 +1,160 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Tests for USP attention layout helpers.""" + +from __future__ import annotations + +import types + +from absl.testing import absltest +import jax + +from maxtext.common.common_types import MODEL_MODE_PREFILL +from maxtext.common.common_types import MODEL_MODE_TRAIN +from maxtext.kernels.attention import ulysses_attention +from maxtext.kernels.attention import usp_attention + + +class UspAttentionTest(absltest.TestCase): + + def test_context_parallel_strategy_helper_identifies_usp(self): + self.assertTrue( + usp_attention.is_context_parallel_usp_requested(types.SimpleNamespace(context_parallel_strategy="usp")) + ) + self.assertFalse( + usp_attention.is_context_parallel_usp_requested(types.SimpleNamespace(context_parallel_strategy="ulysses")) + ) + + def test_validate_usp_runtime_allows_train_mode(self): + usp_attention.validate_usp_runtime(model_mode=MODEL_MODE_TRAIN) + + def test_validate_usp_runtime_rejects_unsupported_runtime_features(self): + with self.assertRaisesRegex(ValueError, "train mode"): + usp_attention.validate_usp_runtime(model_mode=MODEL_MODE_PREFILL) + with self.assertRaisesRegex(ValueError, "ragged attention"): + usp_attention.validate_usp_runtime(model_mode=MODEL_MODE_TRAIN, use_ragged_attention=True) + with self.assertRaisesRegex(ValueError, "chunked prefill"): + usp_attention.validate_usp_runtime(model_mode=MODEL_MODE_TRAIN, previous_chunk=object()) + with self.assertRaisesRegex(ValueError, "attention sinks"): + usp_attention.validate_usp_runtime(model_mode=MODEL_MODE_TRAIN, sinks=object()) + with self.assertRaisesRegex(ValueError, "indexer"): + usp_attention.validate_usp_runtime(model_mode=MODEL_MODE_TRAIN, indexer_mask=object()) + with self.assertRaisesRegex(ValueError, "bidirectional"): + usp_attention.validate_usp_runtime(model_mode=MODEL_MODE_TRAIN, bidirectional_mask=object()) + with self.assertRaisesRegex(NotImplementedError, "record_max_logits"): + usp_attention.validate_usp_runtime(model_mode=MODEL_MODE_TRAIN, record_max_logits=True) + + def test_with_usp_sequence_axes_preserves_partition_spec_type(self): + spec = jax.sharding.PartitionSpec("data", None, None, "tensor") + + out = usp_attention.with_usp_sequence_axes(spec, "context", "context_ulysses", sequence_dim=2) + + self.assertIsInstance(out, jax.sharding.PartitionSpec) + self.assertEqual(tuple(out), ("data", None, ("context", "context_ulysses"), "tensor")) + + def test_layout_validators_reject_invalid_shardings(self): + mesh = types.SimpleNamespace(shape={"context": 2, "context_ulysses": 2, "tensor": 2}) + pair = ("context", "context_ulysses") + cases = [ + ( + "sequence sharding dimension", + lambda: usp_attention.with_usp_sequence_axes((None, None), *pair, sequence_dim=2), + ), + ( + "unsharded or exactly", + lambda: usp_attention.with_usp_sequence_axes((None, None, "context", None), *pair, sequence_dim=2), + ), + ( + "to differ", + lambda: usp_attention.validate_usp_mesh_axes( + axis_names_q=(None, None, pair, None), + axis_names_kv=(None, None, pair, None), + sequence_dim_q=2, + sequence_dim_kv=2, + mesh=mesh, + ring_axis="context", + ulysses_axis="context", + ), + ), + ( + "mesh axis 'context_ulysses' to exist", + lambda: usp_attention.validate_usp_mesh_axes( + axis_names_q=(None, None, pair, None), + axis_names_kv=(None, None, pair, None), + sequence_dim_q=2, + sequence_dim_kv=2, + mesh=types.SimpleNamespace(shape={"context": 2}), + ring_axis="context", + ulysses_axis="context_ulysses", + ), + ), + ( + "only on the sequence dimension", + lambda: usp_attention.validate_usp_mesh_axes( + axis_names_q=(None, "context_ulysses", pair, None), + axis_names_kv=(None, None, pair, None), + sequence_dim_q=2, + sequence_dim_kv=2, + mesh=mesh, + ring_axis="context", + ulysses_axis="context_ulysses", + ), + ), + ( + "Q sequence sharding to be exactly", + lambda: usp_attention.validate_usp_mesh_axes( + axis_names_q=(None, None, "context", None), + axis_names_kv=(None, None, pair, None), + sequence_dim_q=2, + sequence_dim_kv=2, + mesh=mesh, + ring_axis="context", + ulysses_axis="context_ulysses", + ), + ), + ( + "K/V sequence sharding to be exactly", + lambda: usp_attention.validate_usp_mesh_axes( + axis_names_q=(None, None, pair, None), + axis_names_kv=(None, None, None, None), + sequence_dim_q=2, + sequence_dim_kv=2, + mesh=mesh, + ring_axis="context", + ulysses_axis="context_ulysses", + ), + ), + ( + "TPU USP attention requires local query heads", + lambda: ulysses_attention.validate_head_sharding( + axis_names_q=(None, "tensor", pair, None), + axis_names_kv=(None, "tensor", pair, None), + mesh=mesh, + num_query_heads=4, + num_kv_heads=4, + head_dim_q=1, + head_dim_kv=1, + ulysses_size=4, + attention_label="TPU USP attention", + ), + ), + ] + for expected_regex, invoke in cases: + with self.subTest(expected_regex=expected_regex): + with self.assertRaisesRegex(ValueError, expected_regex): + invoke() + + +if __name__ == "__main__": + absltest.main() diff --git a/tests/unit/usp_collective_test.py b/tests/unit/usp_collective_test.py new file mode 100644 index 0000000000..bf4308efc6 --- /dev/null +++ b/tests/unit/usp_collective_test.py @@ -0,0 +1,170 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Executes the USP collectives on a forced multi-device CPU mesh. + +Runs as a subprocess so the forced device count takes effect before JAX +initializes; the parent pytest process has already initialized JAX with the +default device count. The child checks, against a dense single-device +reference, the Ulysses exchange and round trip over the pair sharding, the +segment-ID gather into the ring layout, per-ring-chunk attention equivalence, +and independent Q, K, and V gradients, on 2x2, 2x4, and 4x2 ring x Ulysses +meshes and on a 3-D fsdp x ring x Ulysses mesh with the batch sharded over +fsdp. The ring dimension is modeled by gathering K/V over the ring axis; the +rotation itself belongs to the TPU ring kernel tests. The asymmetric +factorizations matter: a symmetric mesh cannot expose an accidental swap of +the ring and Ulysses dimensions. +""" + +import os +import subprocess +import sys +from functools import partial + +import jax +import jax.numpy as jnp +from jax.sharding import Mesh, PartitionSpec as P +import numpy as np +import pytest + +from maxtext.kernels.attention import ulysses_attention + + +@pytest.mark.cpu_only +def test_usp_collectives_match_dense_reference_on_cpu_mesh(): + env = os.environ.copy() + env["XLA_FLAGS"] = env.get("XLA_FLAGS", "") + " --xla_force_host_platform_device_count=8" + env["JAX_PLATFORMS"] = "cpu" + result = subprocess.run([sys.executable, __file__], env=env, capture_output=True, text=True, check=False) + assert result.returncode == 0, f"stdout:\n{result.stdout}\nstderr:\n{result.stderr}" + assert "USP_COLLECTIVE_CHECKS_PASSED" in result.stdout + + +def _dense_reference_attention(query, key, value, segment_ids): + """Causal segment-masked GQA attention computed on one device.""" + _, num_query_heads, seq_len, _ = query.shape + num_kv_heads = key.shape[1] + group_size = num_query_heads // num_kv_heads + key = jnp.repeat(key, group_size, axis=1) + value = jnp.repeat(value, group_size, axis=1) + + logits = jnp.einsum("bhqd,bhkd->bhqk", query, key) + causal = jnp.tril(jnp.ones((seq_len, seq_len), dtype=bool)) + same_segment = segment_ids[:, :, None] == segment_ids[:, None, :] + not_padding = segment_ids != 0 + mask = causal[None, None, :, :] & same_segment[:, None, :, :] & not_padding[:, None, None, :] + logits = jnp.where(mask, logits, -1e30) + weights = jnp.exp(logits - jnp.max(logits, axis=-1, keepdims=True)) + weights = weights * mask + weights = weights / jnp.maximum(jnp.sum(weights, axis=-1, keepdims=True), 1e-30) + return jnp.einsum("bhqk,bhkd->bhqd", weights, value) + + +def _run_collective_checks(mesh, batch_axis): + """Runs the exchange, segment-ID layout, attention, and gradient checks on one mesh.""" + batch, num_query_heads, num_kv_heads, seq_len, head_dim = 2, 8, 4, 32, 4 + ulysses_axis = "context_ulysses" + ring_axis = "context" + ring_size = mesh.shape[ring_axis] + data_spec = P(batch_axis, None, ("context", "context_ulysses"), None) + segment_spec = P(batch_axis, ("context", "context_ulysses")) + + # Rank-coded values make any head or sequence misordering visible exactly. + def coded(num_heads, offset): + values = np.arange(batch * num_heads * seq_len * head_dim, dtype=np.float32) + return jnp.asarray(values.reshape(batch, num_heads, seq_len, head_dim) / 100.0 + offset) + + query = coded(num_query_heads, 1.0) + key = coded(num_kv_heads, 2.0) + value = coded(num_kv_heads, 3.0) + # Segments begin and end inside different shards of the pair sharding, with + # trailing padding zeros. + segment_ids = jnp.broadcast_to(jnp.asarray([1] * 10 + [2] * 12 + [0] * 10, dtype=jnp.int32)[None, :], (batch, seq_len)) + + # Round trip through the real helpers is exact on the 2-D mesh. + @partial( + jax.shard_map, + mesh=mesh, + in_specs=data_spec, + out_specs=data_spec, + check_vma=False, + ) + def round_trip(tensor): + return ulysses_attention.inverse_ulysses_all_to_all( + ulysses_attention.ulysses_all_to_all(tensor, ulysses_axis), ulysses_axis + ) + + np.testing.assert_array_equal(jax.device_get(round_trip(query)), jax.device_get(query)) + + # The forward exchange produces each rank's head subset over its ring chunk. + @partial( + jax.shard_map, + mesh=mesh, + in_specs=data_spec, + out_specs=P(batch_axis, "context_ulysses", "context", None), + check_vma=False, + ) + def exchange(tensor): + return ulysses_attention.ulysses_all_to_all(tensor, ulysses_axis) + + np.testing.assert_array_equal(jax.device_get(exchange(query)), jax.device_get(query)) + + @partial( + jax.shard_map, + mesh=mesh, + in_specs=(data_spec, data_spec, data_spec, segment_spec), + out_specs=data_spec, + check_vma=False, + ) + def usp_attention_fn(query, key, value, segment_ids): + query = ulysses_attention.ulysses_all_to_all(query, ulysses_axis) + key = ulysses_attention.ulysses_all_to_all(key, ulysses_axis) + value = ulysses_attention.ulysses_all_to_all(value, ulysses_axis) + ring_segment_ids = jax.lax.all_gather(segment_ids, ulysses_axis, axis=1, tiled=True) + # Model the ring dimension by gathering K/V over the ring axis: each device + # computes its own ring chunk of query rows against the full sequence. + full_key = jax.lax.all_gather(key, ring_axis, axis=2, tiled=True) + full_value = jax.lax.all_gather(value, ring_axis, axis=2, tiled=True) + full_segment_ids = jax.lax.all_gather(ring_segment_ids, ring_axis, axis=1, tiled=True) + full_query = jax.lax.all_gather(query, ring_axis, axis=2, tiled=True) + output = _dense_reference_attention(full_query, full_key, full_value, full_segment_ids) + chunk = output.shape[2] // ring_size + output = jax.lax.dynamic_slice_in_dim(output, jax.lax.axis_index(ring_axis) * chunk, chunk, axis=2) + return ulysses_attention.inverse_ulysses_all_to_all(output, ulysses_axis) + + def dense_loss(query, key, value): + output = _dense_reference_attention(query, key, value, segment_ids) + return jnp.sum(output * jnp.cos(output)) + + def usp_loss(query, key, value): + output = usp_attention_fn(query, key, value, segment_ids) + return jnp.sum(output * jnp.cos(output)) + + dense_output = _dense_reference_attention(query, key, value, segment_ids) + usp_output = usp_attention_fn(query, key, value, segment_ids) + np.testing.assert_allclose(jax.device_get(usp_output), jax.device_get(dense_output), atol=1e-5) + + dense_grads = jax.grad(dense_loss, argnums=(0, 1, 2))(query, key, value) + usp_grads = jax.grad(usp_loss, argnums=(0, 1, 2))(query, key, value) + for name, dense_grad, usp_grad in zip(("dQ", "dK", "dV"), dense_grads, usp_grads): + np.testing.assert_allclose(jax.device_get(usp_grad), jax.device_get(dense_grad), atol=1e-5, err_msg=name) + + +if __name__ == "__main__": + _devices = np.array(jax.devices()) + assert len(_devices) == 8, jax.devices() + _run_collective_checks(Mesh(_devices[:4].reshape(2, 2), ("context", "context_ulysses")), batch_axis=None) + _run_collective_checks(Mesh(_devices.reshape(2, 4), ("context", "context_ulysses")), batch_axis=None) + _run_collective_checks(Mesh(_devices.reshape(4, 2), ("context", "context_ulysses")), batch_axis=None) + _run_collective_checks(Mesh(_devices.reshape((2, 2, 2)), ("fsdp", "context", "context_ulysses")), batch_axis="fsdp") + print("USP_COLLECTIVE_CHECKS_PASSED") diff --git a/tests/utils/sharding_info/deepseek2-16b/tpu7x-16/slice_1/rule_default/named_shardings.json b/tests/utils/sharding_info/deepseek2-16b/tpu7x-16/slice_1/rule_default/named_shardings.json index b535731554..17acc46352 100644 --- a/tests/utils/sharding_info/deepseek2-16b/tpu7x-16/slice_1/rule_default/named_shardings.json +++ b/tests/utils/sharding_info/deepseek2-16b/tpu7x-16/slice_1/rule_default/named_shardings.json @@ -8,6 +8,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -21,6 +22,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -40,6 +42,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -53,6 +56,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -76,6 +80,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -89,6 +94,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -100,6 +106,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -125,6 +132,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -138,6 +146,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -149,6 +158,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -174,6 +184,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -187,6 +198,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -205,6 +217,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -223,6 +236,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -236,6 +250,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -261,6 +276,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -274,6 +290,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -299,6 +316,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -312,6 +330,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -337,6 +356,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -350,6 +370,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -369,6 +390,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -388,6 +410,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -401,6 +424,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -413,6 +437,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -439,6 +464,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -452,6 +478,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -464,6 +491,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -484,6 +512,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -497,6 +526,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -509,6 +539,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -535,6 +566,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -548,6 +580,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -560,6 +593,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -582,6 +616,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -595,6 +630,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -606,7 +642,8 @@ [ "fsdp", "fsdp_transpose", - "context" + "context", + "context_ulysses" ], null, null @@ -626,6 +663,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -639,6 +677,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -651,7 +690,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -676,6 +716,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -689,6 +730,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -701,7 +743,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -726,6 +769,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -739,6 +783,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -757,7 +802,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -776,6 +822,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -789,6 +836,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -800,6 +848,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -825,6 +874,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -838,6 +888,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -849,6 +900,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -874,6 +926,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -887,6 +940,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -905,6 +959,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -923,6 +978,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -936,6 +992,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -961,6 +1018,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -974,6 +1032,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -999,6 +1058,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1012,6 +1072,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1037,6 +1098,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1050,6 +1112,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1069,6 +1132,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1088,6 +1152,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1101,6 +1166,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1113,6 +1179,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1139,6 +1206,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1152,6 +1220,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1164,6 +1233,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1184,6 +1254,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1197,6 +1268,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1209,6 +1281,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1235,6 +1308,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1248,6 +1322,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1265,6 +1340,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1282,6 +1358,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1295,6 +1372,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1314,6 +1392,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1327,6 +1406,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1350,6 +1430,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1363,6 +1444,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1374,6 +1456,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -1399,6 +1482,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1412,6 +1496,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1423,6 +1508,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -1448,6 +1534,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1461,6 +1548,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1479,6 +1567,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -1497,6 +1586,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1510,6 +1600,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1535,6 +1626,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1548,6 +1640,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1573,6 +1666,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1586,6 +1680,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1611,6 +1706,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1624,6 +1720,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1643,6 +1740,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1662,6 +1760,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1675,6 +1774,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1687,6 +1787,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1713,6 +1814,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1726,6 +1828,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1738,6 +1841,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1758,6 +1862,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1771,6 +1876,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1783,6 +1889,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1809,6 +1916,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1822,6 +1930,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1834,6 +1943,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -1856,6 +1966,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1869,6 +1980,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1880,7 +1992,8 @@ [ "fsdp", "fsdp_transpose", - "context" + "context", + "context_ulysses" ], null, null @@ -1900,6 +2013,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1913,6 +2027,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1925,7 +2040,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -1950,6 +2066,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1963,6 +2080,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1975,7 +2093,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -2000,6 +2119,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2013,6 +2133,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2031,7 +2152,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -2050,6 +2172,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2063,6 +2186,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2074,6 +2198,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -2099,6 +2224,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2112,6 +2238,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2123,6 +2250,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -2148,6 +2276,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2161,6 +2290,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2179,6 +2309,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -2197,6 +2328,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2210,6 +2342,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2235,6 +2368,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2248,6 +2382,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2273,6 +2408,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2286,6 +2422,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2311,6 +2448,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2324,6 +2462,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2343,6 +2482,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2362,6 +2502,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2375,6 +2516,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2387,6 +2529,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -2413,6 +2556,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2426,6 +2570,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2438,6 +2583,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -2458,6 +2604,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2471,6 +2618,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2483,6 +2631,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -2509,6 +2658,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2522,6 +2672,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2539,6 +2690,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2556,6 +2708,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2569,6 +2722,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2592,6 +2746,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2605,6 +2760,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2616,6 +2772,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -2641,6 +2798,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2654,6 +2812,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2665,6 +2824,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -2690,6 +2850,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2703,6 +2864,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2721,6 +2883,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -2739,6 +2902,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2752,6 +2916,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2777,6 +2942,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2790,6 +2956,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2815,6 +2982,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2828,6 +2996,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2853,6 +3022,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2866,6 +3036,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2885,6 +3056,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2904,6 +3076,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2917,6 +3090,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2929,6 +3103,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -2955,6 +3130,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2968,6 +3144,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2980,6 +3157,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3000,6 +3178,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3013,6 +3192,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3025,6 +3205,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3051,6 +3232,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3064,6 +3246,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3076,6 +3259,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -3098,6 +3282,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3111,6 +3296,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3122,7 +3308,8 @@ [ "fsdp", "fsdp_transpose", - "context" + "context", + "context_ulysses" ], null, null @@ -3142,6 +3329,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3155,6 +3343,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3167,7 +3356,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -3192,6 +3382,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3205,6 +3396,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3217,7 +3409,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -3242,6 +3435,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3255,6 +3449,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3273,7 +3468,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -3292,6 +3488,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3305,6 +3502,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3316,6 +3514,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -3341,6 +3540,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3354,6 +3554,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3365,6 +3566,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -3390,6 +3592,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3403,6 +3606,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3421,6 +3625,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -3439,6 +3644,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3452,6 +3658,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3477,6 +3684,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3490,6 +3698,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3515,6 +3724,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3528,6 +3738,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3553,6 +3764,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3566,6 +3778,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3585,6 +3798,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -3604,6 +3818,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3617,6 +3832,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3629,6 +3845,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3655,6 +3872,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3668,6 +3886,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3680,6 +3899,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3700,6 +3920,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3713,6 +3934,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3725,6 +3947,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3751,6 +3974,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3764,6 +3988,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3781,6 +4006,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -3798,6 +4024,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3811,6 +4038,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, diff --git a/tests/utils/sharding_info/deepseek2-16b/v6e-16/slice_1/rule_default_ici_fsdp_parallelism=-1_ici_expert_parallelism=4/named_shardings.json b/tests/utils/sharding_info/deepseek2-16b/v6e-16/slice_1/rule_default_ici_fsdp_parallelism=-1_ici_expert_parallelism=4/named_shardings.json index c6144b3769..12ac73ac76 100644 --- a/tests/utils/sharding_info/deepseek2-16b/v6e-16/slice_1/rule_default_ici_fsdp_parallelism=-1_ici_expert_parallelism=4/named_shardings.json +++ b/tests/utils/sharding_info/deepseek2-16b/v6e-16/slice_1/rule_default_ici_fsdp_parallelism=-1_ici_expert_parallelism=4/named_shardings.json @@ -8,6 +8,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -21,6 +22,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -40,6 +42,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -53,6 +56,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -76,6 +80,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -89,6 +94,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -100,6 +106,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -125,6 +132,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -138,6 +146,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -149,6 +158,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -174,6 +184,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -187,6 +198,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -205,6 +217,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -223,6 +236,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -236,6 +250,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -261,6 +276,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -274,6 +290,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -299,6 +316,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -312,6 +330,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -337,6 +356,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -350,6 +370,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -369,6 +390,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -388,6 +410,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -401,6 +424,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -413,6 +437,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -439,6 +464,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -452,6 +478,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -464,6 +491,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -484,6 +512,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -497,6 +526,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -509,6 +539,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -535,6 +566,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -548,6 +580,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -560,6 +593,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -582,6 +616,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -595,6 +630,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -606,7 +642,8 @@ [ "fsdp", "fsdp_transpose", - "context" + "context", + "context_ulysses" ], null, null @@ -626,6 +663,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -639,6 +677,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -651,7 +690,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -676,6 +716,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -689,6 +730,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -701,7 +743,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -726,6 +769,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -739,6 +783,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -757,7 +802,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -776,6 +822,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -789,6 +836,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -800,6 +848,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -825,6 +874,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -838,6 +888,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -849,6 +900,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -874,6 +926,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -887,6 +940,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -905,6 +959,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -923,6 +978,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -936,6 +992,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -961,6 +1018,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -974,6 +1032,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -999,6 +1058,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1012,6 +1072,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1037,6 +1098,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1050,6 +1112,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1069,6 +1132,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1088,6 +1152,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1101,6 +1166,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1113,6 +1179,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1139,6 +1206,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1152,6 +1220,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1164,6 +1233,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1184,6 +1254,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1197,6 +1268,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1209,6 +1281,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1235,6 +1308,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1248,6 +1322,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1265,6 +1340,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1282,6 +1358,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1295,6 +1372,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1314,6 +1392,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1327,6 +1406,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1350,6 +1430,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1363,6 +1444,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1374,6 +1456,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -1399,6 +1482,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1412,6 +1496,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1423,6 +1508,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -1448,6 +1534,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1461,6 +1548,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1479,6 +1567,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -1497,6 +1586,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1510,6 +1600,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1535,6 +1626,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1548,6 +1640,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1573,6 +1666,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1586,6 +1680,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1611,6 +1706,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1624,6 +1720,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1643,6 +1740,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1662,6 +1760,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1675,6 +1774,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1687,6 +1787,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1713,6 +1814,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1726,6 +1828,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1738,6 +1841,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1758,6 +1862,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1771,6 +1876,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1783,6 +1889,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -1809,6 +1916,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1822,6 +1930,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1834,6 +1943,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -1856,6 +1966,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1869,6 +1980,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1880,7 +1992,8 @@ [ "fsdp", "fsdp_transpose", - "context" + "context", + "context_ulysses" ], null, null @@ -1900,6 +2013,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1913,6 +2027,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1925,7 +2040,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -1950,6 +2066,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1963,6 +2080,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1975,7 +2093,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -2000,6 +2119,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2013,6 +2133,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2031,7 +2152,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -2050,6 +2172,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2063,6 +2186,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2074,6 +2198,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -2099,6 +2224,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2112,6 +2238,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2123,6 +2250,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -2148,6 +2276,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2161,6 +2290,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2179,6 +2309,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -2197,6 +2328,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2210,6 +2342,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2235,6 +2368,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2248,6 +2382,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2273,6 +2408,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2286,6 +2422,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2311,6 +2448,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2324,6 +2462,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2343,6 +2482,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2362,6 +2502,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2375,6 +2516,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2387,6 +2529,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -2413,6 +2556,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2426,6 +2570,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2438,6 +2583,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -2458,6 +2604,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2471,6 +2618,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2483,6 +2631,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -2509,6 +2658,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2522,6 +2672,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2539,6 +2690,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2556,6 +2708,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2569,6 +2722,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2592,6 +2746,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2605,6 +2760,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2616,6 +2772,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -2641,6 +2798,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2654,6 +2812,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2665,6 +2824,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -2690,6 +2850,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2703,6 +2864,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2721,6 +2883,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -2739,6 +2902,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2752,6 +2916,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2777,6 +2942,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2790,6 +2956,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2815,6 +2982,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2828,6 +2996,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2853,6 +3022,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2866,6 +3036,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2885,6 +3056,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2904,6 +3076,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2917,6 +3090,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2929,6 +3103,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -2955,6 +3130,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2968,6 +3144,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2980,6 +3157,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3000,6 +3178,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3013,6 +3192,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3025,6 +3205,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3051,6 +3232,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3064,6 +3246,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3076,6 +3259,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -3098,6 +3282,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3111,6 +3296,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3122,7 +3308,8 @@ [ "fsdp", "fsdp_transpose", - "context" + "context", + "context_ulysses" ], null, null @@ -3142,6 +3329,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3155,6 +3343,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3167,7 +3356,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -3192,6 +3382,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3205,6 +3396,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3217,7 +3409,8 @@ null, [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -3242,6 +3435,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3255,6 +3449,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3273,7 +3468,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -3292,6 +3488,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3305,6 +3502,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3316,6 +3514,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -3341,6 +3540,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3354,6 +3554,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3365,6 +3566,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], null, @@ -3390,6 +3592,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3403,6 +3606,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3421,6 +3625,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -3439,6 +3644,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3452,6 +3658,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3477,6 +3684,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3490,6 +3698,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3515,6 +3724,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3528,6 +3738,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3553,6 +3764,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3566,6 +3778,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3585,6 +3798,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -3604,6 +3818,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3617,6 +3832,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3629,6 +3845,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3655,6 +3872,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3668,6 +3886,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3680,6 +3899,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3700,6 +3920,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3713,6 +3934,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3725,6 +3947,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], null, @@ -3751,6 +3974,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3764,6 +3988,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3781,6 +4006,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -3798,6 +4024,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3811,6 +4038,7 @@ "fsdp": 4, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, diff --git a/tests/utils/sharding_info/gpt-oss-20b/tpu7x-16/slice_1/rule_default/named_shardings.json b/tests/utils/sharding_info/gpt-oss-20b/tpu7x-16/slice_1/rule_default/named_shardings.json index 970aa1d4ef..6e4f43ed36 100644 --- a/tests/utils/sharding_info/gpt-oss-20b/tpu7x-16/slice_1/rule_default/named_shardings.json +++ b/tests/utils/sharding_info/gpt-oss-20b/tpu7x-16/slice_1/rule_default/named_shardings.json @@ -8,6 +8,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -21,6 +22,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -40,6 +42,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -53,6 +56,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -76,6 +80,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -89,6 +94,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -120,6 +126,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -133,6 +140,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -145,6 +153,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -171,6 +180,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -184,6 +194,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -196,6 +207,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -214,6 +226,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -227,6 +240,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -246,6 +260,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -265,6 +280,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -278,6 +294,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -309,6 +326,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -322,6 +340,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -334,6 +353,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -360,6 +380,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -373,6 +394,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -398,6 +420,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -411,6 +434,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -442,6 +466,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -455,6 +480,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -467,6 +493,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -493,6 +520,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -506,6 +534,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -531,6 +560,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -544,6 +574,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -556,6 +587,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -576,6 +608,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -589,6 +622,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -601,7 +635,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -626,6 +661,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -639,6 +675,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -669,6 +706,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -682,6 +720,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -694,7 +733,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -719,6 +759,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -732,6 +773,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -762,6 +804,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -775,6 +818,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -793,7 +837,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -812,6 +857,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -825,6 +871,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -852,6 +899,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -865,6 +913,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -890,6 +939,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -903,6 +953,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -928,6 +979,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -941,6 +993,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -972,6 +1025,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -985,6 +1039,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -997,6 +1052,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1023,6 +1079,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1036,6 +1093,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1048,6 +1106,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -1066,6 +1125,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1079,6 +1139,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1098,6 +1159,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1117,6 +1179,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1130,6 +1193,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1161,6 +1225,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1174,6 +1239,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1186,6 +1252,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1212,6 +1279,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1225,6 +1293,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1250,6 +1319,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1263,6 +1333,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1294,6 +1365,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1307,6 +1379,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1319,6 +1392,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1345,6 +1419,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1358,6 +1433,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1383,6 +1459,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1396,6 +1473,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1408,6 +1486,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1428,6 +1507,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1441,6 +1521,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1453,7 +1534,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -1478,6 +1560,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1491,6 +1574,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1521,6 +1605,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1534,6 +1619,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1546,7 +1632,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -1571,6 +1658,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1584,6 +1672,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1614,6 +1703,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1627,6 +1717,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1645,7 +1736,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -1664,6 +1756,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1677,6 +1770,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1704,6 +1798,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1717,6 +1812,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1742,6 +1838,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1755,6 +1852,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1780,6 +1878,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1793,6 +1892,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1805,6 +1905,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -1827,6 +1928,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1840,6 +1942,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1857,6 +1960,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1874,6 +1978,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1887,6 +1992,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1906,6 +2012,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1919,6 +2026,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1942,6 +2050,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1955,6 +2064,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1986,6 +2096,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1999,6 +2110,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2011,6 +2123,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2037,6 +2150,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2050,6 +2164,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2062,6 +2177,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -2080,6 +2196,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2093,6 +2210,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2112,6 +2230,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2131,6 +2250,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2144,6 +2264,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2175,6 +2296,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2188,6 +2310,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2200,6 +2323,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2226,6 +2350,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2239,6 +2364,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2264,6 +2390,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2277,6 +2404,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2308,6 +2436,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2321,6 +2450,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2333,6 +2463,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2359,6 +2490,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2372,6 +2504,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2397,6 +2530,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2410,6 +2544,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2422,6 +2557,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2442,6 +2578,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2455,6 +2592,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2467,7 +2605,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -2492,6 +2631,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2505,6 +2645,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2535,6 +2676,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2548,6 +2690,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2560,7 +2703,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -2585,6 +2729,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2598,6 +2743,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2628,6 +2774,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2641,6 +2788,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2659,7 +2807,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -2678,6 +2827,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2691,6 +2841,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2718,6 +2869,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2731,6 +2883,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2756,6 +2909,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2769,6 +2923,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2794,6 +2949,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2807,6 +2963,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2838,6 +2995,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2851,6 +3009,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2863,6 +3022,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2889,6 +3049,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2902,6 +3063,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2914,6 +3076,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -2932,6 +3095,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2945,6 +3109,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2964,6 +3129,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2983,6 +3149,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2996,6 +3163,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3027,6 +3195,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3040,6 +3209,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3052,6 +3222,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -3078,6 +3249,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3091,6 +3263,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3116,6 +3289,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3129,6 +3303,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3160,6 +3335,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3173,6 +3349,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3185,6 +3362,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -3211,6 +3389,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3224,6 +3403,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3249,6 +3429,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3262,6 +3443,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3274,6 +3456,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -3294,6 +3477,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3307,6 +3491,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3319,7 +3504,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -3344,6 +3530,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3357,6 +3544,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3387,6 +3575,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3400,6 +3589,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3412,7 +3602,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -3437,6 +3628,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3450,6 +3642,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3480,6 +3673,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3493,6 +3687,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3511,7 +3706,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -3530,6 +3726,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3543,6 +3740,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3570,6 +3768,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3583,6 +3782,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3608,6 +3808,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3621,6 +3822,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3646,6 +3848,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3659,6 +3862,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3671,6 +3875,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -3693,6 +3898,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3706,6 +3912,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3723,6 +3930,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -3740,6 +3948,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3753,6 +3962,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3776,6 +3986,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3789,6 +4000,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3820,6 +4032,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3833,6 +4046,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3845,6 +4059,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -3871,6 +4086,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3884,6 +4100,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3896,6 +4113,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -3914,6 +4132,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3927,6 +4146,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3946,6 +4166,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -3965,6 +4186,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3978,6 +4200,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4009,6 +4232,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4022,6 +4246,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4034,6 +4259,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4060,6 +4286,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4073,6 +4300,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4098,6 +4326,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4111,6 +4340,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4142,6 +4372,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4155,6 +4386,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4167,6 +4399,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4193,6 +4426,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4206,6 +4440,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4231,6 +4466,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4244,6 +4480,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4256,6 +4493,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4276,6 +4514,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4289,6 +4528,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4301,7 +4541,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -4326,6 +4567,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4339,6 +4581,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4369,6 +4612,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4382,6 +4626,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4394,7 +4639,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -4419,6 +4665,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4432,6 +4679,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4462,6 +4710,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4475,6 +4724,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4493,7 +4743,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -4512,6 +4763,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4525,6 +4777,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4552,6 +4805,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4565,6 +4819,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4590,6 +4845,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4603,6 +4859,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4628,6 +4885,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4641,6 +4899,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4672,6 +4931,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4685,6 +4945,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4697,6 +4958,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4723,6 +4985,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4736,6 +4999,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4748,6 +5012,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -4766,6 +5031,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4779,6 +5045,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4798,6 +5065,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -4817,6 +5085,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4830,6 +5099,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4861,6 +5131,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4874,6 +5145,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4886,6 +5158,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4912,6 +5185,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4925,6 +5199,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4950,6 +5225,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4963,6 +5239,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4994,6 +5271,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5007,6 +5285,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5019,6 +5298,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -5045,6 +5325,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5058,6 +5339,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5083,6 +5365,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5096,6 +5379,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5108,6 +5392,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -5128,6 +5413,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5141,6 +5427,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5153,7 +5440,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -5178,6 +5466,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5191,6 +5480,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5221,6 +5511,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5234,6 +5525,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5246,7 +5538,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -5271,6 +5564,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5284,6 +5578,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5314,6 +5609,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5327,6 +5623,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5345,7 +5642,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -5364,6 +5662,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5377,6 +5676,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5404,6 +5704,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5417,6 +5718,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5442,6 +5744,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5455,6 +5758,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5480,6 +5784,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5493,6 +5798,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5505,6 +5811,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -5527,6 +5834,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5540,6 +5848,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5557,6 +5866,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -5574,6 +5884,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5587,6 +5898,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, diff --git a/tests/utils/sharding_info/gpt-oss-20b/tpu7x-16/slice_1/rule_default_ici_fsdp_parallelism=-1_ici_expert_parallelism=2/named_shardings.json b/tests/utils/sharding_info/gpt-oss-20b/tpu7x-16/slice_1/rule_default_ici_fsdp_parallelism=-1_ici_expert_parallelism=2/named_shardings.json index 1a90c3e31f..05e1324f3d 100644 --- a/tests/utils/sharding_info/gpt-oss-20b/tpu7x-16/slice_1/rule_default_ici_fsdp_parallelism=-1_ici_expert_parallelism=2/named_shardings.json +++ b/tests/utils/sharding_info/gpt-oss-20b/tpu7x-16/slice_1/rule_default_ici_fsdp_parallelism=-1_ici_expert_parallelism=2/named_shardings.json @@ -8,6 +8,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -21,6 +22,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -40,6 +42,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -53,6 +56,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -76,6 +80,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -89,6 +94,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -120,6 +126,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -133,6 +140,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -145,6 +153,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -171,6 +180,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -184,6 +194,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -196,6 +207,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -214,6 +226,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -227,6 +240,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -246,6 +260,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -265,6 +280,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -278,6 +294,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -309,6 +326,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -322,6 +340,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -334,6 +353,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -360,6 +380,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -373,6 +394,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -398,6 +420,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -411,6 +434,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -442,6 +466,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -455,6 +480,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -467,6 +493,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -493,6 +520,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -506,6 +534,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -531,6 +560,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -544,6 +574,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -556,6 +587,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -576,6 +608,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -589,6 +622,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -601,7 +635,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -626,6 +661,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -639,6 +675,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -669,6 +706,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -682,6 +720,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -694,7 +733,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -719,6 +759,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -732,6 +773,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -762,6 +804,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -775,6 +818,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -793,7 +837,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -812,6 +857,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -825,6 +871,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -852,6 +899,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -865,6 +913,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -890,6 +939,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -903,6 +953,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -928,6 +979,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -941,6 +993,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -972,6 +1025,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -985,6 +1039,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -997,6 +1052,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1023,6 +1079,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1036,6 +1093,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1048,6 +1106,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -1066,6 +1125,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1079,6 +1139,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1098,6 +1159,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1117,6 +1179,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1130,6 +1193,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1161,6 +1225,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1174,6 +1239,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1186,6 +1252,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1212,6 +1279,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1225,6 +1293,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1250,6 +1319,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1263,6 +1333,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1294,6 +1365,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1307,6 +1379,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1319,6 +1392,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1345,6 +1419,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1358,6 +1433,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1383,6 +1459,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1396,6 +1473,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1408,6 +1486,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1428,6 +1507,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1441,6 +1521,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1453,7 +1534,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -1478,6 +1560,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1491,6 +1574,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1521,6 +1605,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1534,6 +1619,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1546,7 +1632,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -1571,6 +1658,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1584,6 +1672,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1614,6 +1703,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1627,6 +1717,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1645,7 +1736,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -1664,6 +1756,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1677,6 +1770,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1704,6 +1798,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1717,6 +1812,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1742,6 +1838,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1755,6 +1852,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1780,6 +1878,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1793,6 +1892,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1805,6 +1905,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -1827,6 +1928,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1840,6 +1942,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1857,6 +1960,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1874,6 +1978,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1887,6 +1992,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1906,6 +2012,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1919,6 +2026,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1942,6 +2050,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1955,6 +2064,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1986,6 +2096,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1999,6 +2110,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2011,6 +2123,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2037,6 +2150,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2050,6 +2164,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2062,6 +2177,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -2080,6 +2196,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2093,6 +2210,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2112,6 +2230,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2131,6 +2250,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2144,6 +2264,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2175,6 +2296,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2188,6 +2310,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2200,6 +2323,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2226,6 +2350,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2239,6 +2364,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2264,6 +2390,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2277,6 +2404,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2308,6 +2436,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2321,6 +2450,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2333,6 +2463,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2359,6 +2490,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2372,6 +2504,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2397,6 +2530,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2410,6 +2544,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2422,6 +2557,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2442,6 +2578,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2455,6 +2592,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2467,7 +2605,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -2492,6 +2631,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2505,6 +2645,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2535,6 +2676,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2548,6 +2690,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2560,7 +2703,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -2585,6 +2729,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2598,6 +2743,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2628,6 +2774,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2641,6 +2788,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2659,7 +2807,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -2678,6 +2827,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2691,6 +2841,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2718,6 +2869,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2731,6 +2883,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2756,6 +2909,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2769,6 +2923,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2794,6 +2949,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2807,6 +2963,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2838,6 +2995,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2851,6 +3009,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2863,6 +3022,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -2889,6 +3049,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2902,6 +3063,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2914,6 +3076,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -2932,6 +3095,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2945,6 +3109,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -2964,6 +3129,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -2983,6 +3149,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -2996,6 +3163,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3027,6 +3195,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3040,6 +3209,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3052,6 +3222,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -3078,6 +3249,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3091,6 +3263,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3116,6 +3289,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3129,6 +3303,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3160,6 +3335,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3173,6 +3349,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3185,6 +3362,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -3211,6 +3389,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3224,6 +3403,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3249,6 +3429,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3262,6 +3443,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3274,6 +3456,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -3294,6 +3477,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3307,6 +3491,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3319,7 +3504,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -3344,6 +3530,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3357,6 +3544,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3387,6 +3575,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3400,6 +3589,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3412,7 +3602,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -3437,6 +3628,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3450,6 +3642,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3480,6 +3673,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3493,6 +3687,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3511,7 +3706,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -3530,6 +3726,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3543,6 +3740,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3570,6 +3768,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3583,6 +3782,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3608,6 +3808,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3621,6 +3822,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3646,6 +3848,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3659,6 +3862,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3671,6 +3875,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -3693,6 +3898,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3706,6 +3912,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3723,6 +3930,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -3740,6 +3948,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3753,6 +3962,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3776,6 +3986,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3789,6 +4000,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3820,6 +4032,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3833,6 +4046,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3845,6 +4059,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -3871,6 +4086,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3884,6 +4100,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3896,6 +4113,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -3914,6 +4132,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3927,6 +4146,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -3946,6 +4166,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -3965,6 +4186,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -3978,6 +4200,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4009,6 +4232,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4022,6 +4246,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4034,6 +4259,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4060,6 +4286,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4073,6 +4300,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4098,6 +4326,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4111,6 +4340,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4142,6 +4372,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4155,6 +4386,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4167,6 +4399,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4193,6 +4426,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4206,6 +4440,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4231,6 +4466,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4244,6 +4480,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4256,6 +4493,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4276,6 +4514,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4289,6 +4528,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4301,7 +4541,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -4326,6 +4567,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4339,6 +4581,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4369,6 +4612,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4382,6 +4626,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4394,7 +4639,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -4419,6 +4665,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4432,6 +4679,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4462,6 +4710,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4475,6 +4724,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4493,7 +4743,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -4512,6 +4763,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4525,6 +4777,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4552,6 +4805,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4565,6 +4819,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4590,6 +4845,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4603,6 +4859,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4628,6 +4885,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4641,6 +4899,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4672,6 +4931,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4685,6 +4945,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4697,6 +4958,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4723,6 +4985,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4736,6 +4999,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4748,6 +5012,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage" @@ -4766,6 +5031,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4779,6 +5045,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4798,6 +5065,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -4817,6 +5085,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4830,6 +5099,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4861,6 +5131,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4874,6 +5145,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4886,6 +5158,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -4912,6 +5185,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4925,6 +5199,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4950,6 +5225,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -4963,6 +5239,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -4994,6 +5271,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5007,6 +5285,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5019,6 +5298,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -5045,6 +5325,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5058,6 +5339,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5083,6 +5365,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5096,6 +5379,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5108,6 +5392,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -5128,6 +5413,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5141,6 +5427,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5153,7 +5440,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -5178,6 +5466,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5191,6 +5480,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5221,6 +5511,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5234,6 +5525,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5246,7 +5538,8 @@ "stage", [ "fsdp", - "context" + "context", + "context_ulysses" ], [ "fsdp_transpose", @@ -5271,6 +5564,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5284,6 +5578,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5314,6 +5609,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5327,6 +5623,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5345,7 +5642,8 @@ ], [ "fsdp", - "context" + "context", + "context_ulysses" ] ], "shape": [ @@ -5364,6 +5662,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5377,6 +5676,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5404,6 +5704,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5417,6 +5718,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5442,6 +5744,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5455,6 +5758,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5480,6 +5784,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5493,6 +5798,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5505,6 +5811,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], [ @@ -5527,6 +5834,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5540,6 +5848,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -5557,6 +5866,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -5574,6 +5884,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -5587,6 +5898,7 @@ "fsdp": 8, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, diff --git a/tests/utils/sharding_info/qwen3-0.6b/tpu7x-16/slice_1/rule_default/named_shardings.json b/tests/utils/sharding_info/qwen3-0.6b/tpu7x-16/slice_1/rule_default/named_shardings.json index 3c56796fba..0019ed198e 100644 --- a/tests/utils/sharding_info/qwen3-0.6b/tpu7x-16/slice_1/rule_default/named_shardings.json +++ b/tests/utils/sharding_info/qwen3-0.6b/tpu7x-16/slice_1/rule_default/named_shardings.json @@ -8,6 +8,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -21,6 +22,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -40,6 +42,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -53,6 +56,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -76,6 +80,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -89,6 +94,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -100,6 +106,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], "stage", @@ -125,6 +132,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -138,6 +146,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -149,6 +158,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], "stage", @@ -174,6 +184,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -187,6 +198,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -205,6 +217,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -223,6 +236,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -236,6 +250,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -261,6 +276,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -274,6 +290,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -299,6 +316,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -312,6 +330,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -324,6 +343,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -350,6 +370,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -363,6 +384,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -388,6 +410,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -401,6 +424,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -420,6 +444,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -439,6 +464,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -452,6 +478,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -464,6 +491,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -490,6 +518,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -503,6 +532,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -528,6 +558,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -541,6 +572,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -553,6 +585,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -579,6 +612,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -592,6 +626,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -609,6 +644,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -626,6 +662,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -639,6 +676,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -658,6 +696,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -671,6 +710,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -694,6 +734,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -707,6 +748,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -718,6 +760,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], "stage", @@ -743,6 +786,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -756,6 +800,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -767,6 +812,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], "stage", @@ -792,6 +838,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -805,6 +852,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -823,6 +871,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -841,6 +890,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -854,6 +904,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -879,6 +930,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -892,6 +944,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -917,6 +970,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -930,6 +984,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -942,6 +997,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -968,6 +1024,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -981,6 +1038,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1006,6 +1064,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1019,6 +1078,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1038,6 +1098,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1057,6 +1118,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1070,6 +1132,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1082,6 +1145,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1108,6 +1172,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1121,6 +1186,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1146,6 +1212,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1159,6 +1226,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1171,6 +1239,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1197,6 +1266,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1210,6 +1280,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1227,6 +1298,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1244,6 +1316,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1257,6 +1330,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1280,6 +1354,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1293,6 +1368,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1304,6 +1380,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], "stage", @@ -1329,6 +1406,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1342,6 +1420,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1353,6 +1432,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ], "stage", @@ -1378,6 +1458,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1391,6 +1472,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1409,6 +1491,7 @@ [ "fsdp", "context", + "context_ulysses", "expert" ] ], @@ -1427,6 +1510,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1440,6 +1524,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1465,6 +1550,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1478,6 +1564,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1503,6 +1590,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1516,6 +1604,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1528,6 +1617,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1554,6 +1644,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1567,6 +1658,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1592,6 +1684,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1605,6 +1698,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1624,6 +1718,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1643,6 +1738,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1656,6 +1752,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1668,6 +1765,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1694,6 +1792,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1707,6 +1806,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1732,6 +1832,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1745,6 +1846,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1757,6 +1859,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ], "stage", @@ -1783,6 +1886,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1796,6 +1900,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1, @@ -1813,6 +1918,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "expert" ] ], @@ -1830,6 +1936,7 @@ "fsdp", "fsdp_transpose", "context", + "context_ulysses", "context_autoregressive", "tensor", "tensor_sequence", @@ -1843,6 +1950,7 @@ "fsdp": 16, "fsdp_transpose": 1, "context": 1, + "context_ulysses": 1, "context_autoregressive": 1, "tensor": 1, "tensor_sequence": 1,