Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
59 commits
Select commit Hold shift + click to select a range
84f81ea
chore: update development base image to 26.09-py3
svcnemo-autobot Sep 29, 2026
1ff951b
fix: update DeepGEMM for PyTorch C++20 requirements
svcnemo-autobot Sep 29, 2026
2560c5a
fix: install elfutils headers for DeepGEMM builds
svcnemo-autobot Sep 29, 2026
7640398
build: backport Mamba C++20 flags for PyTorch 26.09
balasaajay Sep 30, 2026
dac0cff
build: fix DeepGEMM FP8 headers on CUDA 12.9
balasaajay Sep 30, 2026
d58a15d
fix(ci): avoid failing permission prepass on CoreWeave runners
balasaajay Sep 30, 2026
46d3cca
fix(tests): restore MoE benchmark RNG and failure propagation
balasaajay Sep 30, 2026
67ee88c
fix(inference): honor FA4 split-KV and attention-sink requirements
balasaajay Sep 30, 2026
6ebe61d
build: align CuTe and QuACK dependencies with FlashAttention 4
balasaajay Sep 30, 2026
e9f01a0
build: satisfy FA4 TVM FFI requirements with compatible TileLang
balasaajay Sep 30, 2026
2f16cb1
fix: backport FA4 packed subtraction compatibility
balasaajay Sep 30, 2026
71f91fc
test: preserve historical MoE benchmark routing
balasaajay Sep 30, 2026
64dd800
test: refresh NanoV3 GB200 batch128 performance baseline
balasaajay Sep 30, 2026
0e2c866
test: refresh H100 FSDP context-parallel loss reference
balasaajay Sep 30, 2026
e1e0c85
fix: handle missing FSDP checkpoint version metadata
balasaajay Sep 30, 2026
87a35cd
test: refresh stable H100 DeepSeek loss references
balasaajay Sep 30, 2026
a883731
test: refresh stable H100 GPT performance improvement
balasaajay Sep 30, 2026
85f8fc5
test: refresh H100 DeepSeek overlap loss reference
balasaajay Sep 30, 2026
b028d4e
test: refresh H100 hybrid FSDP loss reference
balasaajay Sep 30, 2026
28f8046
test: refresh H100 FSDP v2 overlap loss reference
balasaajay Sep 30, 2026
e071ec3
test: refresh H100 HSDP loss reference
balasaajay Sep 30, 2026
74ad5bb
test: refresh H100 MoE GRPO timing reference
balasaajay Sep 30, 2026
070712a
test: refresh H100 pipelined GPT timing reference
balasaajay Sep 30, 2026
20838da
test: refresh stable T5 references at full precision
balasaajay Sep 30, 2026
4906863
build: backport NCCL EP window offsets for PyTorch 26.09
balasaajay Sep 30, 2026
9b8300d
test: clear inference mode after batch-invariance tests
balasaajay Sep 30, 2026
c9819a9
test: refresh stable A100 T5 reference at full precision
balasaajay Sep 30, 2026
bd4fe43
fix: register CUDA graph generators on PyTorch 26.09
balasaajay Sep 30, 2026
ee07a6b
fix: remove unused Torch version import
balasaajay Sep 30, 2026
513efe4
fix: select compatible paged attention for small Hopper heads
balasaajay Sep 30, 2026
45679b9
fix: match native clamp boundary gradients in fused SwiGLU
balasaajay Sep 30, 2026
39bb2bc
fix: preserve token-only padding during SSM decode
balasaajay Sep 30, 2026
718410b
test: dump thread stacks during stalled pytest teardown
balasaajay Sep 30, 2026
0b07373
test: retain every rank stderr in CI console logs
balasaajay Sep 30, 2026
6fef589
test: restore existing tests for the image update
balasaajay Sep 30, 2026
16c1b9a
Merge main and retain Transformer Engine 2.20
balasaajay Oct 1, 2026
12b1c0a
fix: guard Hopper attention dispatch by CUDA device
balasaajay Oct 1, 2026
3befeee
test: disable dev cases that leak inference mode
balasaajay Oct 1, 2026
8f99759
fix: defer clamp boundary probe until backward
balasaajay Oct 1, 2026
4f4d214
Merge branch 'main' into bump-dev-image-26.09-pr7687-20260930
balasaajay Oct 2, 2026
3e3e659
test: quarantine A100 LTS MoE reference mismatch
balasaajay Oct 3, 2026
88879c2
test: cover PyTorch 26.09 kernel compatibility fixes
balasaajay Oct 3, 2026
4546569
test: refresh goldens within the one-percent numerical limit
balasaajay Oct 6, 2026
ec34db6
ci: capture transformer shutdown diagnostics on H100
balasaajay Oct 6, 2026
4114504
ci: avoid blocking NCCL finalization in transformer validation
balasaajay Oct 6, 2026
c5378c3
Merge main and retain the TE 2.20 branch update
balasaajay Oct 6, 2026
b73d6cb
fix(ci): retain native MIMO packages and isolate 26.09 regressions
balasaajay Oct 6, 2026
25b5f63
ci: quarantine two NCCL-crashing transformer variants
balasaajay Oct 6, 2026
64977df
Merge branch 'main' into bump-dev-image-26.09-pr7687-20260930
balasaajay Oct 6, 2026
d842b7a
ci: validate PyTorch 26.09 with FlashAttention 2
balasaajay Oct 7, 2026
090200d
Merge branch 'main' into bump-dev-image-26.09-pr7687-20260930
balasaajay Oct 7, 2026
fab15a7
fix: detect CUDA graph generator registration at runtime
ksivaman Oct 7, 2026
a7a6534
test: re-enable six GPT CP2 GitHub cases
balasaajay Oct 7, 2026
23e74cc
test: pin generic inference fixtures to FlashAttention 2
balasaajay Oct 8, 2026
896ea0e
test: refresh repeatable FA2 loss and gradient-zero references
balasaajay Oct 8, 2026
46e7128
test: pin remaining generic inference fixtures to FlashAttention 2
balasaajay Oct 8, 2026
880404c
test: pin text generation controller fixtures to FlashAttention 2
balasaajay Oct 8, 2026
89fde7c
test: quarantine H100 distributed optimizer reference mismatches
balasaajay Oct 8, 2026
5b38ce0
test: pin custom functional launchers to FlashAttention 2
balasaajay Oct 8, 2026
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
4 changes: 2 additions & 2 deletions .gitlab/stages/01.build.yml
Original file line number Diff line number Diff line change
Expand Up @@ -62,12 +62,12 @@ test:pre_build_image:
- IMAGE: CI_MCORE_DEV_IMAGE
FILE: Dockerfile.ci.dev
IMAGE_TYPE: dev
BASE_IMAGE: nvcr.io/nvidia/pytorch:26.08-py3
BASE_IMAGE: nvcr.io/nvidia/pytorch:26.09-py3
PLATFORM: amd64
- IMAGE: CI_MCORE_DEV_IMAGE
FILE: Dockerfile.ci.dev
IMAGE_TYPE: dev
BASE_IMAGE: nvcr.io/nvidia/pytorch:26.08-py3
BASE_IMAGE: nvcr.io/nvidia/pytorch:26.09-py3
PLATFORM: arm64
- IMAGE: UTILITY_IMAGE
FILE: Dockerfile.linting
Expand Down
2 changes: 1 addition & 1 deletion docker/.ngc_version.dev
Original file line number Diff line number Diff line change
@@ -1 +1 @@
nvcr.io/nvidia/pytorch:26.08-py3
nvcr.io/nvidia/pytorch:26.09-py3
24 changes: 23 additions & 1 deletion docker/Dockerfile.ci.dev
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,8 @@ ENV UV_LINK_MODE=copy

RUN bash -ex <<"EOF"
apt-get update
apt-get install -y --no-install-recommends gettext python3-venv psmisc uuid-runtime
# DeepGEMM/DeepJIT requires elfutils/libdwfl.h for its exception support.
apt-get install -y --no-install-recommends gettext python3-venv psmisc uuid-runtime libdw-dev
apt-get clean
python -m venv /opt/jet
ARCH=$(uname -m)
Expand Down Expand Up @@ -88,6 +89,7 @@ RUN --mount=type=bind,source=docker/common/install_nccl.sh,target=/opt/install_n
EOF

COPY pyproject.toml uv.lock /workspace/
COPY docker/patches/mamba-cxx20.patch /workspace/mamba-cxx20.patch
COPY megatron/core/__init__.py /workspace/megatron/core/
COPY megatron/core/package_info.py /workspace/megatron/core/
ARG IMAGE_TYPE=dev
Expand Down Expand Up @@ -131,6 +133,7 @@ RUN --mount=type=cache,id=uv-${TARGETARCH}-${NVTE_CUDA_ARCHS},target=/root/.cach
--no-install-package torch \
--no-install-package torchvision \
--no-install-package triton \
--no-install-package mamba-ssm \
--no-install-package transformer-engine-cu12 \
--no-install-package nvidia-cublas-cu12 \
--no-install-package nvidia-cuda-cupti-cu12 \
Expand Down Expand Up @@ -161,6 +164,25 @@ RUN --mount=type=cache,id=uv-${TARGETARCH}-${NVTE_CUDA_ARCHS},target=/root/.cach
cp -a "${NCCL_EP_JIT_HEADERS}/../nccl_ep.h" /opt/nccl-ep/include/
test -f /opt/nccl-ep/include/nccl_ep/device/ht_ep.cuh
test -f /opt/nccl-ep/include/nccl_ep.h

# Backport state-spaces/mamba#1000 without changing the pinned runtime sources.
# PyTorch's ATen headers now require C++20 for both host and CUDA compilation.
MAMBA_REV=$(sed -n 's/^mamba-ssm = .*rev = "\([^"]*\)".*/\1/p' /workspace/pyproject.toml)
test -n "${MAMBA_REV}"
MAMBA_BUILD_DIR=$(mktemp -d)
trap 'rm -rf "${MAMBA_BUILD_DIR}"' EXIT
git -C "${MAMBA_BUILD_DIR}" init
git -C "${MAMBA_BUILD_DIR}" fetch --depth 1 https://github.com/state-spaces/mamba.git "${MAMBA_REV}"
git -C "${MAMBA_BUILD_DIR}" checkout --detach "${MAMBA_REV}"
git -C "${MAMBA_BUILD_DIR}" apply /workspace/mamba-cxx20.patch
MAMBA_FORCE_BUILD=TRUE uv pip install --no-deps --no-build-isolation -v "${MAMBA_BUILD_DIR}"
# Keep the source revision and applied patch after deleting the build checkout.
mkdir -p /opt/mamba-build-info
git -C "${MAMBA_BUILD_DIR}" rev-parse HEAD > /opt/mamba-build-info/revision
cp /workspace/mamba-cxx20.patch /opt/mamba-build-info/
rm -rf "${MAMBA_BUILD_DIR}"
trap - EXIT

EOF

# Reuse the wheel without copying a second archive into the final image.
Expand Down
3 changes: 2 additions & 1 deletion docker/common/install.sh
Original file line number Diff line number Diff line change
Expand Up @@ -86,7 +86,8 @@ main() {

# Install tools
apt-get update
apt-get install -y wget curl git cmake
# DeepGEMM/DeepJIT requires elfutils/libdwfl.h for its exception support.
apt-get install -y wget curl git cmake libdw-dev

# Install CUDA
if [[ "$BASE_IMAGE" == "ubuntu" ]]; then
Expand Down
33 changes: 33 additions & 0 deletions docker/patches/mamba-cxx20.patch
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
Backport of https://github.com/state-spaces/mamba/commit/653923ce8fb0d47cdd9bcfd5904a0f1d58f91274
Build the existing pinned Mamba sources with the standard required by PyTorch's ATen headers.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same here.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This patch fixes native-extension build failures because our pinned Mamba explicitly uses C++17 while the new PyTorch ATen headers require C++20; it changes only four compiler flags. No published Mamba release currently contains the fix, and state-spaces/mamba#1000 remains unmerged, so we retain the patch until a suitable upstream revision is adopted and validated.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can you please add this comment to the actual file?


diff --git a/setup.py b/setup.py
--- a/setup.py
+++ b/setup.py
@@ -210,10 +210,10 @@ def append_nvcc_threads(nvcc_extra_args):
if HIP_BUILD:

extra_compile_args = {
- "cxx": ["-O3", "-std=c++17"],
+ "cxx": ["-O3", "-std=c++20"],
"nvcc": [
"-O3",
- "-std=c++17",
+ "-std=c++20",
f"--offload-arch={os.getenv('HIP_ARCHITECTURES', 'native')}",
"-U__CUDA_NO_HALF_OPERATORS__",
"-U__CUDA_NO_HALF_CONVERSIONS__",
@@ -223,11 +223,11 @@ def append_nvcc_threads(nvcc_extra_args):
}
else:
extra_compile_args = {
- "cxx": ["-O3", "-std=c++17"],
+ "cxx": ["-O3", "-std=c++20"],
"nvcc": append_nvcc_threads(
[
"-O3",
- "-std=c++17",
+ "-std=c++20",
"-U__CUDA_NO_HALF_OPERATORS__",
"-U__CUDA_NO_HALF_CONVERSIONS__",
"-U__CUDA_NO_BFLOAT16_OPERATORS__",
40 changes: 38 additions & 2 deletions megatron/core/fusions/fused_bias_swiglu.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,14 +3,44 @@

# pylint: disable=missing-function-docstring, missing-class-docstring

from functools import cache
from typing import Optional

import torch
import torch.nn.functional as F
from torch._subclasses.fake_tensor import unset_fake_temporarily

from megatron.core.jit import jit_fuser
from megatron.core.utils import nvtx_decorator


@cache
def _probe_clamp_boundaries():
"""Read the installed scalar-clamp subgradient using only CPU constants."""
# PyTorch changed this within the 2.14 development series, so a version check
# cannot distinguish builds. FunctionsManual.cpp at b2c75dd062 uses strict
# inequalities; 4fdf77b940 uses inclusive ones. Probe both scalar-clamp forms
# once on CPU, including calls under no_grad or inference_mode.
# This zero-input probe needs real CPU constants, even during fake tracing.
with unset_fake_temporarily(), torch.inference_mode(False), torch.enable_grad():
gate = torch.tensor(1.0, dtype=torch.float32, device="cpu", requires_grad=True)
linear = torch.tensor([-1.0, 1.0], dtype=torch.float32, device="cpu", requires_grad=True)
gate_grad, linear_grad = torch.autograd.grad(
gate.clamp(max=1.0) + linear.clamp(min=-1.0, max=1.0).sum(), (gate, linear)
)
gradients = [gate_grad.item(), *linear_grad.tolist()]
if gradients not in ([0.0, 0.0, 0.0], [1.0, 1.0, 1.0]):
raise RuntimeError(f"Unsupported scalar-clamp boundary gradients: {gradients}")
return gradients[0] == 1.0


@torch.compiler.assume_constant_result
def _clamp_includes_boundaries():
# Delay autograd initialization until this clamped backward is actually used.
# Keep the cached probe outside tracing and specialize only its Python bool.
return _probe_clamp_boundaries()


###### BIAS SWIGLU FUSION/ NO AUTOGRAD ################


Expand Down Expand Up @@ -163,14 +193,20 @@ def clamped_swiglu_back(g, y, clamp_value):
y_1, y_2 = torch.chunk(y.to(torch.float32), 2, -1)
y_1c = y_1.clamp(min=None, max=clamp_value)
y_2c = y_2.clamp(min=-clamp_value, max=clamp_value)
if _clamp_includes_boundaries():
gate_mask = y_1 <= clamp_value
linear_mask = (y_2 >= -clamp_value) & (y_2 <= clamp_value)
else:
gate_mask = y_1 < clamp_value
linear_mask = (y_2 > -clamp_value) & (y_2 < clamp_value)
res = torch.cat(
(
g
* torch.sigmoid(y_1c)
* (1 + y_1c * (1 - torch.sigmoid(y_1c)))
* y_2c
* (y_1 <= clamp_value).to(g.dtype),
g * F.silu(y_1c) * ((y_2 >= -clamp_value) & (y_2 <= clamp_value)).to(g.dtype),
* gate_mask.to(g.dtype),
g * F.silu(y_1c) * linear_mask.to(g.dtype),
),
-1,
)
Expand Down
40 changes: 18 additions & 22 deletions megatron/core/ssm/ssm_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,18 +223,15 @@ def ssm_dynamic_inference(
if decode_req_count > 0:
seq_len = 1 + context.num_speculative_tokens
decode_token_count = decode_req_count * seq_len
if context.batch_invariant_mode:
# Batch-invariant execution may include token-only rows to preserve
# model-wide M alignment. Those rows do not represent requests and
# must not be passed to the recurrent decode kernels.
assert decode_token_count <= zxBCdt.shape[0], (
"Batch-invariant SSM metadata describes more decode tokens "
f"({decode_token_count}) than the input projection contains "
f"({zxBCdt.shape[0]})."
)
zxBCdt_decode = zxBCdt[:decode_token_count]
else:
zxBCdt_decode = zxBCdt[:decode_token_count] if prefill_req_count > 0 else zxBCdt
# Token-only padding can come from kernel-level batch invariance even
# when the model/context flag is off. Only metadata-backed requests
# belong in the recurrent decode kernels.
assert decode_token_count <= zxBCdt.shape[0], (
"SSM metadata describes more decode tokens "
f"({decode_token_count}) than the input projection contains "
f"({zxBCdt.shape[0]})."
)
zxBCdt_decode = zxBCdt[:decode_token_count]
# Reshape from [N*S, 1, d] to [N, S, d] for the decode kernels.
zxBCdt_decode = zxBCdt_decode.squeeze(1).view(decode_req_count, seq_len, -1)
y_decode = self.ssm_decode(
Expand Down Expand Up @@ -282,16 +279,15 @@ def ssm_dynamic_inference(
else:
raise RuntimeError("Dynamic inference called with 0 decode and 0 prefill requests")

if context.batch_invariant_mode:
# Restore the projection's token-only padding before the output projection.
# Its row count can be TP-local, unlike the context's global token count.
padding_token_count = zxBCdt.shape[0] - y.shape[0]
assert padding_token_count >= 0, (
"Batch-invariant SSM produced more token rows "
f"({y.shape[0]}) than the input projection contained ({zxBCdt.shape[0]})."
)
if padding_token_count > 0:
y = torch.cat((y, y.new_zeros(padding_token_count, *y.shape[1:])), dim=0)
# Restore the projection's token-only padding before the output projection.
# Its row count can be TP-local, unlike the context's global token count.
padding_token_count = zxBCdt.shape[0] - y.shape[0]
assert padding_token_count >= 0, (
"SSM produced more token rows "
f"({y.shape[0]}) than the input projection contained ({zxBCdt.shape[0]})."
)
if padding_token_count > 0:
y = torch.cat((y, y.new_zeros(padding_token_count, *y.shape[1:])), dim=0)

# Zero padding positions to avoid corrupting quantization amax calculations.
if is_using_quantization_scales(self.config):
Expand Down
41 changes: 39 additions & 2 deletions megatron/core/tensor_parallel/random.py
Original file line number Diff line number Diff line change
Expand Up @@ -448,16 +448,53 @@ def get_all_rng_states():
return {}


_CUDAGRAPH_NEEDS_GENERATOR_REGISTRATION = None


def _probe_cudagraph_needs_generator_registration() -> bool:
"""Capture an RNG op on an unregistered generator and report whether capture rejects it."""
generator = torch.Generator(device="cuda")
graph = torch.cuda.CUDAGraph()
try:
with torch.cuda.graph(graph, capture_error_mode="thread_local"):
torch.rand(1, device="cuda", generator=generator)
except RuntimeError:
# Without lazy registration, capture fails with "RNG op during graph capture but
# generator is not registered with the capturing graph". Treat any capture failure
# as needing registration: on PyTorch with lazy registration, registering anyway
# only costs a deprecation warning, while skipping it on older builds is fatal.
return True
finally:
del graph
return False


def cudagraph_needs_generator_registration() -> bool:
"""Whether generators must be registered with a `torch.cuda.CUDAGraph` before capture.

PyTorch >= 2.14 (pytorch/pytorch#176753) lazily registers every generator whose Philox
PyTorch with pytorch/pytorch#176753 lazily registers every generator whose Philox
state is consumed during capture, and `CUDAGraph.register_generator_state()` became a
deprecated no-op that prints a warning on *every* call. Skip the explicit registration
there: it does nothing, and with one call per layer, per graph and per generator it floods
stderr (tens of thousands of lines per rank for dynamic inference with CUDA graphs).

The version string cannot identify that change: 2.14.0a0 nightlies built before it landed
still require explicit registration. Builds older than 2.14.0a0 always require it; for
newer builds, a small capture probe decides, and its result is cached.
"""
return not is_torch_min_version("2.14.0a0")
global _CUDAGRAPH_NEEDS_GENERATOR_REGISTRATION
if _CUDAGRAPH_NEEDS_GENERATOR_REGISTRATION is not None:
return _CUDAGRAPH_NEEDS_GENERATOR_REGISTRATION
if not is_torch_min_version("2.14.0a0"):
needs_registration = True
elif torch.cuda.is_current_stream_capturing():
# The probe cannot capture while another capture is active; register to be safe
# and probe on a later call.
return True
else:
needs_registration = _probe_cudagraph_needs_generator_registration()
_CUDAGRAPH_NEEDS_GENERATOR_REGISTRATION = needs_registration
return needs_registration


def prime_cuda_rng_states_for_graph_capture() -> None:
Expand Down
6 changes: 3 additions & 3 deletions megatron/training/checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -3033,9 +3033,9 @@ def load_checkpoint(
)
# Same as the torch_dist branch: optimizer load templates select checkpoint-era keys by
# version. fsdp_dtensor state has no dtype-keyed FQNs today; keep both paths identical.
optim_sd_kwargs['metadata']['checkpoint_version'] = (
state_dict.get('checkpoint_version') or 0
)
optim_sd_kwargs['metadata']['checkpoint_version'] = (state_dict or {}).get(
'checkpoint_version'
) or 0

# Megatron-FSDP materializes optimizer slots with a dummy zero-gradient
# step while building a loading state dict. A normal full resume
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -255,7 +255,7 @@ requires-dist = ["torch", "packaging", "ninja"]
flash_mla = [
{ git = "https://github.com/deepseek-ai/FlashMLA", rev = "nv_dev" },
]
deep_gemm = { git = "https://github.com/deepseek-ai/DeepGEMM.git", rev = "714dd1a4a980f7937a74343d19a8eba4fe321480" }
deep_gemm = { git = "https://github.com/deepseek-ai/DeepGEMM.git", rev = "ea2b7805b97956268bed1adf94591f428ded6cc1" }
transformer-engine = { git = "https://github.com/NVIDIA/TransformerEngine.git", rev = "6ea2a74a9e98c99e6d7b164a33775cc457520027" }
nemo-run = { git = "https://github.com/NVIDIA-NeMo/Run.git", rev = "e3935393a290aed1822af52139b4b8ee270fed1f" }
emerging_optimizers = { git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git", rev = "v0.3.0" }
Expand Down
6 changes: 6 additions & 0 deletions tests/functional_tests/shell_test_utils/_run_training.sh
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,9 @@ if [[ "$BEFORE_SCRIPT" != null ]]; then
eval "$BEFORE_SCRIPT"
fi

# Keep functional validation on FA2 while using the unpatched NGC 26.09 stack.
export NVTE_FLASH_ATTN_V4=0

# Exit earlier to leave time for properly saving checkpoint
if [[ "$IS_NEMO_TEST" == "true" ]]; then
PARAMS=()
Expand Down Expand Up @@ -187,6 +190,9 @@ fi

# Extract training params
PARAMS=("${PARAMS[@]}" "${TRAINING_PARAMS_ARRAY[@]}")
if [[ "$IS_NEMO_TEST" != "true" ]]; then
PARAMS+=("--flash-attention-version" "2")
fi

# Set PYTHONPATH
export PYTHONPATH="$(pwd):${PYTHONPATH:-}"
Expand Down
Loading
Loading