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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions benchmarks/vbench/run_tpu_generation.sh
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,8 @@ WAN_OVERRIDES=(
"use_batched_text_encoder|USE_BATCHED_TEXT_ENCODER|true"
"flash_block_sizes|FLASH_BLOCK_SIZES|"
"prompt_file|PROMPT_FILE|./benchmarks/vbench/prompts_110.txt"
# Matches the shipped config default; set USE_FUSED_ROPE_KERNEL=true to test it.
"use_fused_rope_kernel|USE_FUSED_ROPE_KERNEL|false"
)

usage() {
Expand All @@ -70,6 +72,11 @@ Common options:
RUN_NAME Generation run name (default: wan-inference; videos are saved to <RUN_NAME>/videos)
PROMPT_FILE Prompt file path (default: ./benchmarks/vbench/prompts_110.txt)
CONFIG_FILE WAN config file (default: src/maxdiffusion/configs/base_wan_27b.yml)
USE_FUSED_ROPE_KERNEL
Enable the fused RMSNorm+RoPE Pallas producer (default: false).
Only measured on TPU v6e and v7; refused on other platforms.
FUSED_ROPE_BLOCK_S / FUSED_ROPE_HEAD_BLOCK
Optional sequence and head tile overrides for the fused RoPE Pallas producer.
EXTERNAL_DISK Mounted disk root for large local files (default: /mnt/disks/external_disk)
HF_CACHE_ROOT Hugging Face cache root (default: \$EXTERNAL_DISK/hf_cache)
HF_HOME Hugging Face home directory (default: \$HF_CACHE_ROOT)
Expand Down Expand Up @@ -247,6 +254,8 @@ emit_remote_args() {
emit_remote_arg "${var}"
done
emit_remote_arg_if_explicit VENV_DIR
emit_remote_arg_if_explicit FUSED_ROPE_BLOCK_S
emit_remote_arg_if_explicit FUSED_ROPE_HEAD_BLOCK
for item in "${WAN_OVERRIDES[@]}"; do
IFS='|' read -r key var value <<< "${item}"
emit_remote_arg "${var}"
Expand Down Expand Up @@ -375,6 +384,8 @@ run_generation() {
IFS='|' read -r key var value <<< "${item}"
args+=("${key}=${!var}")
done
[[ -n "${FUSED_ROPE_BLOCK_S:-}" ]] && args+=("fused_rope_block_s=${FUSED_ROPE_BLOCK_S}")
[[ -n "${FUSED_ROPE_HEAD_BLOCK:-}" ]] && args+=("fused_rope_head_block=${FUSED_ROPE_HEAD_BLOCK}")
args+=("seed=12345" "base_output_directory=gs://${GCS_BUCKET}")
"${args[@]}"
}
Expand Down
42 changes: 27 additions & 15 deletions end_to_end/tpu/run_wan_fast_inference.sh
Original file line number Diff line number Diff line change
Expand Up @@ -49,8 +49,20 @@
# autodetection
# ACCEL_TYPE / TPU_ACCELERATOR_TYPE
# force the raw accelerator type (e.g. v6e-8, tpu7x-8)
# ATTENTION / ULYSSES_SHARDS / BQ / BKV / BKV_COMPUTE / BKV_COMPUTE_IN / BQ_DKV / VMEM_LIMIT_BYTES
# ATTENTION / USE_K_CENTERING / ULYSSES_SHARDS / BQ / BKV / BKV_COMPUTE / BKV_COMPUTE_IN / BQ_DKV / VMEM_LIMIT_BYTES
# override the per-platform attention recipe
# USE_FUSED_ROPE_KERNEL / FUSED_ROPE_BLOCK_S / FUSED_ROPE_HEAD_BLOCK
# override the fused RMSNorm+RoPE Pallas producer settings.
# The v6e and v7 profiles turn the kernel ON, unlike the
# YAML default (off) and the generic profile (off), because
# it is measurably faster there: at the full 5-PR stack
# (40 steps, 720p/81f, warm AOT, DVFS unpinned) denoise is
# 125.9s vs 131.2s on v6e-8 and 105.1s vs 106.7s on tpu7x-8.
# Its output is equivalent but not bit-identical to the XLA
# producer: on v6e the bf16 trajectory diverges (PSNR ~15 dB
# against the XLA-path video, visually clean); on tpu7x the
# output was bit-identical in that measurement.
# USE_FUSED_ROPE_KERNEL=false restores the XLA producer.
# DP / CP / PER_DEVICE_BATCH / SEED
# override mesh parallelism, per-device batch, or RNG seed (default 12345)
set -euo pipefail
Expand All @@ -73,8 +85,6 @@ export HF_HUB_ENABLE_HF_TRANSFER=1
export JAX_DEFAULT_MATMUL_PRECISION=bfloat16
export TORCHINDUCTOR_FX_GRAPH_CACHE=1
export XLA_PYTHON_CLIENT_MEM_FRACTION=0.95
# Without these the JAX persistent cache silently skips most entries, so the
# "warm" start still recompiles a large part of the graph.
export JAX_PERSISTENT_CACHE_MIN_ENTRY_SIZE_BYTES=-1
export JAX_PERSISTENT_CACHE_MIN_COMPILE_TIME_SECS=0

Expand All @@ -85,14 +95,7 @@ export TORCHINDUCTOR_CACHE_DIR=${TORCHINDUCTOR_CACHE_DIR:-$CACHE_ROOT/torch_comp
mkdir -p "$CACHE_ROOT/jax" "$CACHE_ROOT/aot_wan$MODEL" "$CACHE_ROOT/converted" \
"$OUTPUT_DIR" "$TMPDIR" "$TORCHINDUCTOR_CACHE_DIR"

# ---------------------------------------------------------------------------
# TPU platform detection
# ---------------------------------------------------------------------------
# Preference order: explicit override -> TPU_ACCELERATOR_TYPE env (set by some
# runtimes) -> GCE metadata "accelerator-type" (e.g. "v6e-8") -> the
# ACCELERATOR_TYPE line inside the "tpu-env" metadata blob. Detection is pure
# metadata/env: it must not initialise the TPU, or it would take the device
# before the real process starts.
# Detect TPU generation via env or GCE metadata without initializing the device.
_tpu_metadata() {
curl -s -f -m 2 -H 'Metadata-Flavor: Google' \
"http://metadata.google.internal/computeMetadata/v1/instance/attributes/$1" 2> /dev/null || true
Expand All @@ -102,8 +105,6 @@ _detect_accel_type() {
local t="${TPU_ACCELERATOR_TYPE:-}"
[ -z "$t" ] && t="$(_tpu_metadata accelerator-type)"
[ -z "$t" ] && t="$(_tpu_metadata tpu-env | sed -n "s/^ACCELERATOR_TYPE: *'\([^']*\)'.*/\1/p")"
# Guard against metadata returning an HTML error page. Real names seen in the
# wild: "v6e-8", "v5litepod-8", "tpu7x-8" (v7 reports as tpu7x, not v7x).
case "$t" in
v[0-9]* | tpu[0-9]*) printf '%s' "$t" ;;
*) printf '' ;;
Expand All @@ -112,8 +113,8 @@ _detect_accel_type() {

ACCEL_TYPE=${ACCEL_TYPE:-$(_detect_accel_type)}
if [ -n "$ACCEL_TYPE" ]; then
TPU_GEN="${ACCEL_TYPE%%-*}" # v6e-8 -> v6e
TPU_CHIPS="${ACCEL_TYPE##*-}" # v6e-8 -> 8
TPU_GEN="${ACCEL_TYPE%%-*}"
TPU_CHIPS="${ACCEL_TYPE##*-}"
else
TPU_GEN=""
TPU_CHIPS=""
Expand Down Expand Up @@ -165,6 +166,7 @@ case "$TPU_PROFILE" in
DEFAULT_BATCHED_TE=true
DEFAULT_VAE_CHUNK=1
DEFAULT_VAE_SPATIAL=8
DEFAULT_FUSED_ROPE=true
;;
v7)
PLATFORM_LIBTPU="$V7_LIBTPU"
Expand All @@ -180,6 +182,7 @@ case "$TPU_PROFILE" in
DEFAULT_BATCHED_TE=true
DEFAULT_VAE_CHUNK=1
DEFAULT_VAE_SPATIAL=8
DEFAULT_FUSED_ROPE=true
;;
*)
echo "== warning: unrecognised or non-v6e/v7 accelerator '${ACCEL_TYPE:-unknown}';" \
Expand All @@ -197,6 +200,7 @@ case "$TPU_PROFILE" in
DEFAULT_BATCHED_TE=true
DEFAULT_VAE_CHUNK=1
DEFAULT_VAE_SPATIAL=8
DEFAULT_FUSED_ROPE=false
;;
esac

Expand All @@ -212,6 +216,7 @@ case "$LIBTPU_INIT_ARGS" in
esac

ATTENTION=${ATTENTION:-$DEFAULT_ATTENTION}
USE_K_CENTERING=${USE_K_CENTERING:-auto} # auto: on for non-ring (v6e), off for ring (v7)
ULYSSES_SHARDS=${ULYSSES_SHARDS:-$DEFAULT_U}
BQ=${BQ:-$DEFAULT_BQ}
BKV=${BKV:-$DEFAULT_BKV}
Expand All @@ -223,6 +228,10 @@ COMPILE_TE=${COMPILE_TE:-$DEFAULT_COMPILE_TE}
USE_BATCHED_TE=${USE_BATCHED_TE:-$DEFAULT_BATCHED_TE}
VAE_SPATIAL=${VAE_SPATIAL:-$DEFAULT_VAE_SPATIAL}
VAE_DECODE_CHUNK=${VAE_DECODE_CHUNK:-$DEFAULT_VAE_CHUNK}
USE_FUSED_ROPE_KERNEL=${USE_FUSED_ROPE_KERNEL:-$DEFAULT_FUSED_ROPE}
FUSED_ROPE_ARGS=()
[ -n "${FUSED_ROPE_BLOCK_S:-}" ] && FUSED_ROPE_ARGS+=("fused_rope_block_s=$FUSED_ROPE_BLOCK_S")
[ -n "${FUSED_ROPE_HEAD_BLOCK:-}" ] && FUSED_ROPE_ARGS+=("fused_rope_head_block=$FUSED_ROPE_HEAD_BLOCK")

# Mesh: context parallelism carries the Ulysses shards, data parallelism takes
# whatever chips remain. Defaults to CP=4 / DP=2 on an 8-chip slice. On a slice
Expand Down Expand Up @@ -272,6 +281,7 @@ python src/maxdiffusion/generate_wan.py "$CONFIG" \
aot_cache_dir="$CACHE_ROOT/aot_wan$MODEL" \
converted_weights_dir="$CACHE_ROOT/converted" \
attention="$ATTENTION" \
use_k_centering="$USE_K_CENTERING" \
ulysses_shards="$ULYSSES_SHARDS" \
ici_data_parallelism="$DP" ici_fsdp_parallelism=1 \
ici_context_parallelism="$CP" ici_tensor_parallelism=1 \
Expand All @@ -282,6 +292,8 @@ python src/maxdiffusion/generate_wan.py "$CONFIG" \
vae_weights_dtype=bfloat16 vae_dtype=bfloat16 \
text_encoder_dtype=bfloat16 compile_text_encoder="$COMPILE_TE" use_batched_text_encoder="$USE_BATCHED_TE" \
use_kv_cache=true use_base2_exp=true use_experimental_scheduler=true \
use_fused_rope_kernel="$USE_FUSED_ROPE_KERNEL" \
${FUSED_ROPE_ARGS[@]+"${FUSED_ROPE_ARGS[@]}"} \
fps=16 ${GUIDANCE_ARGS[@]+"${GUIDANCE_ARGS[@]}"} \
seed="${SEED:-12345}" \
flash_block_sizes="$FLASH_BLOCK_SIZES" \
Expand Down
3 changes: 3 additions & 0 deletions end_to_end/tpu/run_wan_stack_tests.sh
Original file line number Diff line number Diff line change
Expand Up @@ -39,5 +39,8 @@ TESTS=(
"$T/converted_weights_cache_test.py"
"$T/wan/wan_transformer_test.py"
"$T/wan/wan_warmup_coverage_test.py"
# Fused RMSNorm+RoPE producer, Wan switches in config (feat/wan-custom-kernels)
"$T/fused_rmsnorm_rope_pallas_test.py"
"$T/wan_runtime_options_test.py"
)
PYTHONPATH="src${PYTHONPATH:+:$PYTHONPATH}" exec "${PYTHON:-python3}" -m pytest -q -rs "${TESTS[@]}" "$@"
1 change: 1 addition & 0 deletions src/maxdiffusion/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@
"schedulers": [],
"tpu_utils": [],
"train_utils": [],
"wan_runtime_options": [],
"utils": [
"OptionalDependencyNotAvailable",
"is_flax_available",
Expand Down
31 changes: 29 additions & 2 deletions src/maxdiffusion/configs/base_wan_14b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ jit_initializers: True

# Set true to load weights from pytorch
from_pt: True
split_head_dim: True
split_head_dim: False
# Sparse VideoGen (SVG) Attention configuration
use_svg_attention: False
svg_spatial_density: 0.25
Expand All @@ -105,7 +105,34 @@ attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cud

use_base2_exp: True
use_experimental_scheduler: True
use_k_centering: "auto" # auto: on for non-ring ulysses_custom_fixed_m* (virtual, no copy), off for ring paths (virtual via k_mean, plus pmean when R>1)
# use_k_centering: "auto" turns K-centering on for the non-ring
# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no
# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a
# pmean across ring shards when R>1). Set True/False to force either path.
use_k_centering: "auto"
# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK.
# Off by default; only measured on tpu7x and v6e. See base_wan_27b.yml.
use_fused_rope_kernel: False
# Fused RoPE kernel grid: sequence block size, and heads per grid step (-1 =
# all). Tuned at the Wan 2.2 27B shard shape; re-measure for this model.
fused_rope_block_s: 1024
fused_rope_head_block: -1
# Wan inference switches. These change the compiled graph, so they are part of
# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since
# this file sets every key, the env vars are ignored, with a log line if they differ.)
# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction,
# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps).
wan_rope_norm_mode: "exact"
# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer.
wan_fuse_qk_prescale: True
# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default.
wan_splash_transpose_out: False
# Apply classifier-free guidance before unpatchify (halves the unpatchify work).
wan_cfg_before_unpatchify: True
# Pre-scale cached cross-attention K by 1/sqrt(d) once per video.
wan_cross_attn_prescale_kv: False
# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32".
wan_rope_accum: "auto"
# For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
# Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder.
Expand Down
29 changes: 28 additions & 1 deletion src/maxdiffusion/configs/base_wan_1_3b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -79,7 +79,7 @@ jit_initializers: True

# Set true to load weights from pytorch
from_pt: True
split_head_dim: True
split_head_dim: False
# Sparse VideoGen (SVG) Attention configuration
use_svg_attention: False
svg_spatial_density: 0.25
Expand All @@ -102,7 +102,34 @@ attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cud

use_base2_exp: True
use_experimental_scheduler: True
# use_k_centering: "auto" turns K-centering on for the non-ring
# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no
# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a
# pmean across ring shards when R>1). Set True/False to force either path.
use_k_centering: "auto"
# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK.
# Off by default; only measured on tpu7x and v6e. See base_wan_27b.yml.
use_fused_rope_kernel: False
# Fused RoPE kernel grid: sequence block size, and heads per grid step (-1 =
# all). Tuned at the Wan 2.2 27B shard shape; re-measure for this model.
fused_rope_block_s: 1024
fused_rope_head_block: -1
# Wan inference switches. These change the compiled graph, so they are part of
# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since
# this file sets every key, the env vars are ignored, with a log line if they differ.)
# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction,
# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps).
wan_rope_norm_mode: "exact"
# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer.
wan_fuse_qk_prescale: True
# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default.
wan_splash_transpose_out: False
# Apply classifier-free guidance before unpatchify (halves the unpatchify work).
wan_cfg_before_unpatchify: True
# Pre-scale cached cross-attention K by 1/sqrt(d) once per video.
wan_cross_attn_prescale_kv: False
# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32".
wan_rope_accum: "auto"
# For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
# Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder.
Expand Down
38 changes: 35 additions & 3 deletions src/maxdiffusion/configs/base_wan_27b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ jit_initializers: True

# Set true to load weights from pytorch
from_pt: True
split_head_dim: True
split_head_dim: False
# Sparse VideoGen (SVG) Attention configuration
use_svg_attention: False
svg_spatial_density: 0.35
Expand Down Expand Up @@ -127,9 +127,41 @@ attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cud
# that does not equal CP raises ValueError instead of being silently ignored.
use_base2_exp: True
use_experimental_scheduler: True
# auto: on for non-ring ulysses_custom_fixed_m* (virtual, no copy), off for ring
# paths (virtual K-centering via k_mean, plus a pmean across ring shards when R>1).
# use_k_centering: "auto" turns K-centering on for the non-ring
# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no
# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a
# pmean across ring shards when R>1). Set True/False to force either path.
use_k_centering: "auto"
# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK.
# In isolation on one v6e chip it is 1.9x-2.1x faster than the XLA producer
# (40 heads, per-shard seq 9450 / 18900). End-to-end (40 steps, 720p/81f, full
# stack, DVFS unpinned) it cuts denoise 131.2s -> 125.9s on v6e-8 and
# 106.7s -> 105.1s on tpu7x-8; run_wan_fast_inference.sh enables it there.
# Off by default: inlined into the model XLA rounds the unfused producer
# differently, so the output is equivalent but not identical (v6e 720p/81f,
# same seed: 53.8 dB PSNR after 1 denoise step, 34.1 dB after 40). Refused on
# platforms whose RoPE rounding is unmeasured; tuned and measured for inference.
use_fused_rope_kernel: False
# Sequence block size for the fused RoPE kernel grid (clamped to 256 in "fused" norm mode).
fused_rope_block_s: 1024
# Heads processed per grid step; -1 means "all heads in one step".
fused_rope_head_block: -1
# Wan inference switches. These change the compiled graph, so they are part of
# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since
# this file sets every key, the env vars are ignored, with a log line if they differ.)
# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction,
# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps).
wan_rope_norm_mode: "exact"
# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer.
wan_fuse_qk_prescale: True
# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default.
wan_splash_transpose_out: False
# Apply classifier-free guidance before unpatchify (halves the unpatchify work).
wan_cfg_before_unpatchify: True
# Pre-scale cached cross-attention K by 1/sqrt(d) once per video.
wan_cross_attn_prescale_kv: False
# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32".
wan_rope_accum: "auto"
# For attention=ulysses_ring*, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
# Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder.
Expand Down
29 changes: 28 additions & 1 deletion src/maxdiffusion/configs/base_wan_animate.yml
Original file line number Diff line number Diff line change
Expand Up @@ -80,11 +80,38 @@ jit_initializers: True

# Set true to load weights from pytorch
from_pt: True
split_head_dim: True
split_head_dim: False
attention: 'flash' # Supported attention: dot_product, flash, tokamax_flash, cudnn_flash_te, ring, tokamax_ring, ulysses, ulysses_custom, ulysses_custom_fixed_m, ulysses_custom_fixed_m_per_q_block, ulysses_ring, ulysses_ring_custom, ulysses_ring_custom_fixed_m, ulysses_ring_custom_fixed_m_per_q_block
use_base2_exp: True
use_experimental_scheduler: True
# use_k_centering: "auto" turns K-centering on for the non-ring
# ulysses_custom_fixed_m* kernels (virtual: q.k_mean folded into the kernel, no
# collective, no HBM copy) and off for ring kernels (virtual via k_mean, plus a
# pmean across ring shards when R>1). Set True/False to force either path.
use_k_centering: "auto"
# Fused RMSNorm+RoPE+head-transpose Pallas producer for self-attention QK.
# Off by default; only measured on tpu7x and v6e. See base_wan_27b.yml.
use_fused_rope_kernel: False
# Fused RoPE kernel grid: sequence block size, and heads per grid step (-1 =
# all). Tuned at the Wan 2.2 27B shard shape; re-measure for this model.
fused_rope_block_s: 1024
fused_rope_head_block: -1
# Wan inference switches. These change the compiled graph, so they are part of
# the AOT cache fingerprint. (They used to be WAN_* environment variables. Since
# this file sets every key, the env vars are ignored, with a log line if they differ.)
# RMSNorm reduction for the fused RoPE producer: "exact" (XLA reduction,
# bit-identical to the jitted XLA producer) or "fused" (in-kernel, within 2x pair-relative bf16 eps).
wan_rope_norm_mode: "exact"
# Fold log2(e) into Q and 1/sqrt(d) into K inside the fused RoPE producer.
wan_fuse_qk_prescale: True
# Have the splash kernel write its output as [B, H, S, D] (instead of [B, H, D, S]), removing the post-kernel swap. Off by default.
wan_splash_transpose_out: False
# Apply classifier-free guidance before unpatchify (halves the unpatchify work).
wan_cfg_before_unpatchify: True
# Pre-scale cached cross-attention K by 1/sqrt(d) once per video.
wan_cross_attn_prescale_kv: False
# Fused RoPE accumulation: "auto" (measured per-platform default), "dtype", "f32".
wan_rope_accum: "auto"
# For attention=ulysses_ring, hidden Ulysses shard count; ring shards are context / this.
ulysses_shards: -1
# Splits Ulysses all-to-all into head-group chunks. The last chunk carries any remainder.
Expand Down
Loading
Loading