[Feature Request] Support dynamic token counts in TP communication/GEMM overlap
I maintainer di solito rispondono entro 2 giorni
Nessuno ha ancora preso questa issue.
Valutazione
- Difficoltà
- 4/5
- Tempo stimato
- 3-5 giorni
- Idoneità per principianti
- 45/100
- Tipo di issue
- Funzionalità
- Chiarezza
- Specificata chiaramente
- Stato di attività
- Attiva
- Ambito
- ai-infra-agents, machine-learning, performance
Direzione di ricerca
The issue is about TP communication/GEMM overlap in Transformer Engine for variable token counts. Start by reading the prototype description and the linked THD layout documentation. Look at the UserBuffer initialization and overlap paths in the codebase, focusing on bulk, pipeline, and ring-exchange communicators. Understand how tensor views and communication counts are currently computed. The goal is to modify these to use an active shape instead of the fixed buffer size, ensuring the changes work for BF16 and maintain compatibility with TP size divisibility.
Scritto dal modello di indicizzazione a partire dal testo della issue.
Descrizione
Problem
Our use case is packed-sequence training with THD-format inputs and dynamically sized microbatches. In TE's THD layout, T is the total number of packed tokens in a microbatch. As the microbatch size and sequence lengths change, the flattened activations used by TP communication/GEMM overlap have shape [T, hidden_size], and T can differ from one microbatch to the next.
The TP-overlap UserBuffer is initialized with a fixed shape. In the v2.11 code on which my prototype is based, overlap paths use the registered buffer size for tensor views and communication counts. As a result, they cannot directly operate on a smaller [T, hidden_size] tensor while retaining an allocation registered for the maximum token count.
I would like to support this workload without allocating and registering new UserBuffers for different microbatch sizes or padding every microbatch to the maximum token count. I have a prototype based on a fixed-capacity allocation and a per-microbatch active shape, described below for design discussion.
This is related to #1303, which asked whether TP overlap supports variable sequence length. This issue adds a concrete workload and an implementation proposal.
Proposed API
# Distributed setup and overlap configuration are omitted here.
# Existing initialization uses the maximum token capacity.
te.initialize_ub(
[max_tokens, hidden_size],
tp_size,
quantization_modes=[te.UserBufferQuantizationMode.NONE],
dtype=torch.bfloat16,
)
# Before processing a microbatch with fewer tokens:
te.set_ub_active_shape([active_tokens, hidden_size])
# Run the model with tensors sized for active_tokens.
The prototype exposes set_ub_active_shape() on the PyTorch UserBuffer manager. It also exposes set_buffer_active_shape() and active/capacity shape getters on individual CommOverlap and CommOverlapP2P communicators. The manager updates its unquantized bulk, pipeline, and ring-exchange communicators; quantized and external UserBuffers are left at their capacity shape.
Implementation approach in the prototype
Keep allocation capacity separate from the active shape. initialize_ub() registers a buffer sized for the maximum token count. Each communicator stores that capacity shape and a separate active shape. set_buffer_active_shape() updates the logical shape and tensor views; it does not allocate, free, or re-register the underlying UserBuffer. The active data occupies a contiguous prefix of the registered allocation.
Use active sizes throughout the overlap paths. For bulk AG/RS, buffer views, local-chunk offsets, copy-size checks, and communication counts use the active element count rather than the registered capacity. For pipeline and ring exchange, the P2P layout rebuilds its per-rank chunk views and byte offsets using the active first dimension. Reduce-scatter retains its extra P2P chunks within the original allocation.
Constrain the initial behavior. The active shape must be two-dimensional and nonzero, fit within the registered capacity, keep the hidden dimension unchanged, and have a first dimension divisible by TP size. The current manager API applies one active shape to its eligible communicators. Participating TP ranks must therefore use the same active token count before the corresponding overlap operations. The prototype targets unquantized, 2-byte/BF16 UserBuffers for bulk, pipeline, and ring-exchange overlap; dynamic FP8/quantized UserBuffers are outside this initial scope.
The prototype is based on TE v2.11 and has not yet been ported or validated against current main.
Questions for maintainers
- Would Transformer Engine maintainers be interested in supporting TP communication/GEMM overlap when the token count varies between microbatches, as in the THD packed-sequence workload above?
- If so, would registering a UserBuffer at maximum capacity and updating its active shape for each microbatch be an acceptable direction for an upstream contribution?
If the overall direction makes sense, I would also appreciate guidance on whether a manager-level update is the right API and whether BF16 bulk, pipeline, and ring-exchange paths are a useful first scope.
I have a prototype implementation and can port it to current main and submit a focused PR if this direction makes sense.
- Lingua principale
- Python
- Stelle
- 3.6k
- Fork
- 844
- Merge medio
- 4g 13h
- PR unite (30g)
- 51
Preparare l'ambiente
- Nessun Dockerfile né file Docker Compose
- Ha un modello di pull request
- Leggi la guida per i contributori
Come iniziare
- Leggi tutta la issue e poi la guida ai contributi del progetto.
- Commenta sulla issue per dire che te ne occupi tu — evita che due persone facciano lo stesso lavoro.
- Fai un fork del repository e lavora su un branch.
- Apri una pull request che faccia riferimento al numero della issue.
Altre issue di NVIDIA/TransformerEngine
-
[Bug] Backend selection picks FA3 for training with head_dim_qk=192 / v_head_dim=128, but FA3 backward cannot run itForse già presa @yuweih205 l’ha presa 28 giorni fa. Apertaattention
Difficoltà 2/5 1-3 ore Idoneità per principianti 85/100
NVIDIA/TransformerEngine#3481 · 4 commenti ·
I maintainer di solito rispondono entro 2 giorni
-
Increase MAX_TENSOR_NUMApertabug
Difficoltà 2/5 1-3 ore Idoneità per principianti 68/100
NVIDIA/TransformerEngine#2189 · 7 commenti · 5 reazioni ·
I maintainer di solito rispondono entro 2 giorni
-
Multi-tensor swizzle kernels fail with "too many resources requested for launch" (missing __launch_bounds__)Forse già presa @ravimajeti l’ha presa 1 giorno fa. Aperta
Difficoltà 2/5 1-3 ore Idoneità per principianti 25/100
NVIDIA/TransformerEngine#3621 ·
I maintainer di solito rispondono entro 2 giorni
-
[PyTorch] Avoid selecting FA4 for deterministic training on SM120Forse già presa Una pull request collegata a questa issue è aperta o già unita. Apertabug
Difficoltà 3/5 1-2 giorni Idoneità per principianti 35/100
NVIDIA/TransformerEngine#3594 ·
I maintainer di solito rispondono entro 2 giorni
-
[Bug] Float8CurrentScaling / Float8BlockScaling: fp8_quant_* QParams ignore use_power_2_scales / use_f32_scales passed to the constructorForse già presa Una pull request collegata a questa issue è aperta o già unita. Apertabug
Difficoltà 3/5 1-2 giorni Idoneità per principianti 75/100
NVIDIA/TransformerEngine#3593 · 2 commenti ·
I maintainer di solito rispondono entro 2 giorni
Tutte le issue di NVIDIA/TransformerEngine
Issue simili
-
bug llm translation
Difficoltà 2/5 1-3 ore Idoneità per principianti 78/100
I maintainer di solito rispondono entro 1 giorno
-
Arkansas 2025 tax is $1.70 high above $100,000 net taxable income ($3,809 + 3.9% rule)Forse già presa @PavelMakarchuk l’ha presa oggi. Aperta
Difficoltà 2/5 1-3 ore Idoneità per principianti 74/100
PolicyEngine/policyengine-us#9828 ·
I maintainer di solito rispondono entro 2 giorni
-
bug
Difficoltà 2/5 1-3 ore Idoneità per principianti 68/100
jellyfin/jellyfin-mpv-shim#800 ·
I maintainer di solito rispondono entro 1 giorno
-
skillfs: one malformed chat-log line aborts the entire skill-usage analysis (skill_usage_from_chat_logs.py)Forse già presa @zjncs l’ha presa oggi. Apertacomponent:skillfs
Difficoltà 2/5 1-3 ore Idoneità per principianti 82/100
agentic-os-org/ANOLISA#6116 · 1 commento ·
I maintainer di solito rispondono entro 1 giorno
-
P4: low query
Difficoltà 2/5 1-3 ore Idoneità per principianti 76/100
jeffknupp/association#336 ·