Skip to content

Commit 7843ee4

Browse files
committed
fix(preflight): use physical batch for VRAM check
Signed-off-by: Aaron Gonzales <aagonzales@nvidia.com>
1 parent 82c79c4 commit 7843ee4

2 files changed

Lines changed: 57 additions & 7 deletions

File tree

‎src/nemo_safe_synthesizer/preflight/checks/environment.py‎

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -270,9 +270,9 @@ def activation_memory_gib(
270270
) -> float:
271271
r"""Rough activation VRAM on one device given micro-batch geometry.
272272
273-
Uses ``training.batch_size`` (HF ``per_device_train_batch_size``), not
274-
``gradient_accumulation_steps``. Matches bf16-ish training tensors at
275-
2 bytes/element:
273+
Uses the resolved physical microbatch (HF
274+
``per_device_train_batch_size``), not the logical batch or accumulation
275+
count. Matches bf16-ish training tensors at 2 bytes/element:
276276
277277
\[
278278
M_\text{act} \approx B \cdot S \cdot H \cdot L \cdot 2\text{ bytes}
@@ -397,11 +397,13 @@ def check(self, ctx: MetadataView, collector: IssueCollector) -> None:
397397
n_params / 1e9,
398398
)
399399

400+
batching = config.training.resolve_batching()
401+
per_device_batch_size = batching.per_device_train_batch_size
400402
seq_len = getattr(ctx.metadata, "max_seq_length", None)
401403
comp = estimate_training_vram_components(
402404
n_params=n_params,
403405
training_cfg=config.training,
404-
batch_size=config.training.batch_size,
406+
batch_size=per_device_batch_size,
405407
seq_len=seq_len,
406408
hidden_size=getattr(autoconfig, "hidden_size", None),
407409
num_hidden_layers=getattr(autoconfig, "num_hidden_layers", None),
@@ -429,7 +431,7 @@ def check(self, ctx: MetadataView, collector: IssueCollector) -> None:
429431
f"~{comp.overhead_gib:.1f} GiB reserved){qualifier} "
430432
f"exceeds available ~{max_free_gib:.1f} GiB "
431433
f"(training.max_vram_fraction={config.training.max_vram_fraction:.2g}). "
432-
f"Per-device batch_size={config.training.batch_size}. {oom_risk}. "
434+
f"Per-device batch_size={per_device_batch_size}. {oom_risk}. "
433435
"This remains an estimate -- attention blocks, adapters, optimizer state, and "
434436
"checkpointing materially affect real usage."
435437
),

‎tests/preflight/test_preflight.py‎

Lines changed: 50 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -126,7 +126,12 @@ def _metadata(autoconfig=None):
126126
return md
127127

128128
@staticmethod
129-
def _estimated_vram_gib(config: SafeSynthesizerParameters, autoconfig: PretrainedConfig) -> float:
129+
def _estimated_vram_gib(
130+
config: SafeSynthesizerParameters,
131+
autoconfig: PretrainedConfig,
132+
*,
133+
batch_size: int | None = None,
134+
) -> float:
130135
from nemo_safe_synthesizer.preflight.checks.environment import (
131136
estimate_base_model_params,
132137
estimate_training_vram_components,
@@ -138,7 +143,7 @@ def _estimated_vram_gib(config: SafeSynthesizerParameters, autoconfig: Pretraine
138143
comp = estimate_training_vram_components(
139144
n_params=n_params,
140145
training_cfg=config.training,
141-
batch_size=config.training.batch_size,
146+
batch_size=batch_size if batch_size is not None else config.training.batch_size,
142147
seq_len=DEFAULT_MAX_SEQ_LENGTH,
143148
hidden_size=getattr(autoconfig, "hidden_size", None),
144149
num_hidden_layers=getattr(autoconfig, "num_hidden_layers", None),
@@ -175,6 +180,49 @@ def test_ample_vram_is_silent(self, default_config):
175180
assert not any(i.code == "low_vram" for i in issues)
176181
assert not any(i.code == "vram_exceeds_capacity" for i in issues)
177182

183+
def test_automatic_accumulation_uses_physical_batch_for_vram_estimate(self, default_config):
184+
"""A physical microbatch cap prevents a false VRAM issue for the logical batch."""
185+
config = default_config.model_copy(deep=True)
186+
config.training.batch_size = 8
187+
config.training.gradient_accumulation_steps = "auto"
188+
config.training.max_physical_batch_size = 4
189+
autoconfig = self._autoconfig()
190+
metadata = self._metadata(autoconfig=autoconfig)
191+
batch_4_gib = self._estimated_vram_gib(config, autoconfig, batch_size=4)
192+
batch_8_gib = self._estimated_vram_gib(config, autoconfig, batch_size=8)
193+
available_gib = (batch_4_gib + batch_8_gib) / 2
194+
fake_props = MagicMock(total_memory=100 * 1024**3)
195+
196+
with (
197+
patch("torch.cuda.is_available", return_value=True),
198+
patch("nemo_safe_synthesizer.llm.utils.get_max_vram", return_value={0: available_gib / 100}),
199+
patch("torch.cuda.get_device_properties", return_value=fake_props),
200+
):
201+
issues = VRAMHeadroomCheck().run(make_ctx(config=config, metadata=metadata))
202+
203+
assert not any(issue.code in {"low_vram", "vram_exceeds_capacity"} for issue in issues)
204+
205+
def test_automatic_accumulation_reports_physical_batch_when_vram_is_low(self, default_config):
206+
"""The VRAM diagnostic reports the capped physical per-device batch."""
207+
config = default_config.model_copy(deep=True)
208+
config.training.batch_size = 8
209+
config.training.gradient_accumulation_steps = "auto"
210+
config.training.max_physical_batch_size = 4
211+
autoconfig = self._autoconfig()
212+
metadata = self._metadata(autoconfig=autoconfig)
213+
batch_4_gib = self._estimated_vram_gib(config, autoconfig, batch_size=4)
214+
fake_props = MagicMock(total_memory=100 * 1024**3)
215+
216+
with (
217+
patch("torch.cuda.is_available", return_value=True),
218+
patch("nemo_safe_synthesizer.llm.utils.get_max_vram", return_value={0: batch_4_gib / 200}),
219+
patch("torch.cuda.get_device_properties", return_value=fake_props),
220+
):
221+
issues = VRAMHeadroomCheck().run(make_ctx(config=config, metadata=metadata))
222+
223+
issue = next(issue for issue in issues if issue.code in {"low_vram", "vram_exceeds_capacity"})
224+
assert "Per-device batch_size=4" in issue.message
225+
178226
def test_absurd_batch_errors(self, default_config):
179227
"""Per-device batch_size far too large must fail preflight."""
180228
default_config.training.batch_size = 100_000

0 commit comments

Comments
 (0)