From f5b97cc6ef51199c7c030c0fa17c6774d1ab3803 Mon Sep 17 00:00:00 2001 From: kiwigitops Date: Wed, 22 Jul 2026 21:06:44 -0400 Subject: [PATCH] Fix FrozenLake shuffle flag parsing --- examples/frozenlake/train_frozenlake.py | 18 +++++++++++++++++- 1 file changed, 17 insertions(+), 1 deletion(-) diff --git a/examples/frozenlake/train_frozenlake.py b/examples/frozenlake/train_frozenlake.py index 840e30b1e..c68a0f5ef 100644 --- a/examples/frozenlake/train_frozenlake.py +++ b/examples/frozenlake/train_frozenlake.py @@ -99,6 +99,22 @@ # %% import argparse + +def _str_to_bool(value): + if isinstance(value, bool): + return value + + normalized_value = value.lower() + if normalized_value in ("true", "t", "1", "yes", "y"): + return True + if normalized_value in ("false", "f", "0", "no", "n"): + return False + + raise argparse.ArgumentTypeError( + "expected one of: true, false, 1, 0, yes, no" + ) + + arg_parser = argparse.ArgumentParser( description="Train FrozenLake on Gemma4-2B (single-host TPU)." ) @@ -147,7 +163,7 @@ # every multi-turn agent step its env without waiting for a previous wave to # drain. Drop only if KV cache saturates or generation throughput regresses. arg_parser.add_argument("--max_concurrency", type=int, default=512) -arg_parser.add_argument("--shuffle_data", type=bool, default=True) +arg_parser.add_argument("--shuffle_data", type=_str_to_bool, default=True) arg_parser.add_argument("--seed", type=int, default=42) arg_parser.add_argument( "--loss_agg_mode", type=str, default="sequence-mean-token-mean"