AI-Hypercomputer / AI-Hypercomputer/maxtext

Rampup batch size settings are not validated in the Pydantic config

オープン
#4,690 コメント 0 件 リアクション 0 件 担当者 1 名 @NuojCheng が担当を希望しています GitHub で見る
主要言語
Python
スター
2.4k
フォーク
607
平均マージ
2日 19時間
マージ済み PR(30日)
158

説明

The legacy config validates the rampup batch size settings, the Pydantic config does not. Invalid combinations are accepted and either switch rampup off with no diagnostic, or ramp to a batch size that was never requested.

`pyconfig_deprecated.validate_rampup_batch_size` (src/maxtext/configs/pyconfig_deprecated.py:203) asserts five things: `per_device_batch_size_start > 0`, `per_device_batch_size_increment > 0`, `global_rampup_samples > 0`, `per_device_batch_size - per_device_batch_size_start > 0`, and that the difference divides by the increment.

`DatasetGeneral` in src/maxtext/configs/types.py declares the same four fields around line 1365 with no validator. The schedule derivation in `MaxTextConfig` around line 3067 is written defensively, `if self.global_batch_size_to_load_increment > 0:` then `if num_increments > 0:`, so when either guard fails the block falls through and leaves `rampup_end_step = 0`.

Replicating that arithmetic for 8 devices, `expansion_factor_real_data=1`, `gradient_accumulation_steps=1`:

| settings | rampup_end_step | result |
| --- | --- | --- |
| `per_device_batch_size=8, start=4, increment=2, samples=500` | 14 | correct |
| `increment=0` | 0 | rampup silently off |
| `global_rampup_samples=0` | 0 | rampup silently off |
| `start=8, per_device_batch_size=4` | 0 | rampup silently off |
| `per_device_batch_size=9, start=4, increment=2` | 14 | ramps to global batch size 64, 72 was requested |

The first three mean a run configured with `enable_rampup_batch_size=True` trains at a constant batch size with nothing in the logs saying rampup never engaged. The last one is worse, `num_increments = diff // increment` truncates, so ramp-up finishes below the configured `per_device_batch_size` and stays there.

The Pydantic config should reject these at parse time the way the legacy path does. I can send a PR adding a `model_validator(mode="after")` on `DatasetGeneral` with unit tests in tests/unit/configs_value_test.py.

コントリビューションガイド

コントリビューションガイドを開く

評価

この issue はまだ評価されていません。

新しい issue をメールで受け取る

初心者向けの GitHub issue を短くまとめたダイジェスト。