AI-Hypercomputer / AI-Hypercomputer/maxtext

Rampup batch size settings are not validated in the Pydantic config

未關閉
#4,690 0 則留言 0 個 reaction 已指派 1 人 已被 @NuojCheng 認領 在 GitHub 檢視
主要語言
Python
星號
2.4k
分支
607
平均合併
2 天 19 小時
30 天內合併 PR
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 摘要。