From 229c3498ae575ab3e1cf6142b5204bff76b1e676 Mon Sep 17 00:00:00 2001 From: Orbax Authors Date: Fri, 4 Sep 2026 17:57:28 -0700 Subject: [PATCH] No public description PiperOrigin-RevId: 976576328 --- export/orbax/export/obm_configs.py | 33 +++++++++++++ export/orbax/export/obm_configs_test.py | 62 +++++++++++++++++++++++++ 2 files changed, 95 insertions(+) diff --git a/export/orbax/export/obm_configs.py b/export/orbax/export/obm_configs.py index f99c2839d3..1dab7808d3 100644 --- a/export/orbax/export/obm_configs.py +++ b/export/orbax/export/obm_configs.py @@ -158,6 +158,24 @@ class PriorityBatchingPolicy(enum.Enum): PRIORITY_AWARE = "priority_aware" +@dataclasses.dataclass(kw_only=True) +class PriorityAwareBatchingOptions: + """Priority aware batch options for the Orbax model. + + This is used to configure the batching behavior when + `priority_batching_policy` is set to `PRIORITY_AWARE`. + + Attributes: + max_queue_depth: The maximum sum of task sizes to enqueue. Must be + non-negative. + enable_task_resplit: If true, the priority aware batch queue will resplit + tasks into smaller subtasks if needed. Default is False. + """ + + max_queue_depth: int = 0 + enable_task_resplit: bool = False + + @dataclasses.dataclass(kw_only=True) class BatchOptions: """Batch options for the Orbax model. @@ -188,6 +206,9 @@ class BatchOptions: scope is LOCAL_TO_PIPELINE and queue selection policy is ROUND_ROBIN. priority_batching_policy: The priority batching policy for the batch scheduler. Default is STRICT_FIFO. + priority_aware_batching_options: The priority aware batch options for the + batch scheduler when `priority_batching_policy` is set to + `PRIORITY_AWARE`. """ batch_component: BatchComponent @@ -208,6 +229,7 @@ class BatchOptions: priority_batching_policy: PriorityBatchingPolicy = ( PriorityBatchingPolicy.STRICT_FIFO ) + priority_aware_batching_options: PriorityAwareBatchingOptions | None = None def _validate_batch_options( self, @@ -348,6 +370,17 @@ def __post_init__(self): is_low_priority_batch_options=True, ) + if ( + self.priority_batching_policy == PriorityBatchingPolicy.PRIORITY_AWARE + and self.priority_aware_batching_options is not None + ): + if self.priority_aware_batching_options.max_queue_depth < 0: + raise ValueError( + "`priority_aware_batching_options.max_queue_depth` must be" + " non-negative. Got:" + f" {self.priority_aware_batching_options.max_queue_depth}" + ) + @dataclasses.dataclass(kw_only=True) class Jax2ObmOptions: diff --git a/export/orbax/export/obm_configs_test.py b/export/orbax/export/obm_configs_test.py index d84622d424..b3b78d540f 100644 --- a/export/orbax/export/obm_configs_test.py +++ b/export/orbax/export/obm_configs_test.py @@ -350,6 +350,68 @@ def test_low_priority_batch_options_raise_error_with_non_positive_max_enqueued_b ), ) + def test_priority_aware_batching_options_default(self): + batch_options = obm_configs.BatchOptions( + batch_component=obm_configs.BatchComponent.MODEL_FUNCTION, + max_batch_size=16, + priority_aware_batching_options=obm_configs.PriorityAwareBatchingOptions(), + ) + assert batch_options.priority_aware_batching_options is not None + self.assertEqual( + batch_options.priority_aware_batching_options.max_queue_depth, 0 + ) + self.assertFalse( + batch_options.priority_aware_batching_options.enable_task_resplit + ) + + def test_priority_aware_batching_options_custom(self): + batch_options = obm_configs.BatchOptions( + batch_component=obm_configs.BatchComponent.MODEL_FUNCTION, + max_batch_size=16, + priority_aware_batching_options=obm_configs.PriorityAwareBatchingOptions( + max_queue_depth=100, + enable_task_resplit=True, + ), + ) + assert batch_options.priority_aware_batching_options is not None + self.assertEqual( + batch_options.priority_aware_batching_options.max_queue_depth, 100 + ) + self.assertTrue( + batch_options.priority_aware_batching_options.enable_task_resplit + ) + + def test_priority_aware_batching_options_raise_error_with_negative_max_queue_depth( + self, + ): + with self.assertRaisesRegex( + ValueError, + r"`priority_aware_batching_options.max_queue_depth` must be" + r" non-negative. Got: -1", + ): + obm_configs.BatchOptions( + batch_component=obm_configs.BatchComponent.MODEL_FUNCTION, + max_batch_size=16, + priority_batching_policy=obm_configs.PriorityBatchingPolicy.PRIORITY_AWARE, + priority_aware_batching_options=obm_configs.PriorityAwareBatchingOptions( + max_queue_depth=-1, + ), + ) + + def test_priority_aware_batching_options_not_validated_when_strict_fifo(self): + batch_options = obm_configs.BatchOptions( + batch_component=obm_configs.BatchComponent.MODEL_FUNCTION, + max_batch_size=16, + priority_batching_policy=obm_configs.PriorityBatchingPolicy.STRICT_FIFO, + priority_aware_batching_options=obm_configs.PriorityAwareBatchingOptions( + max_queue_depth=-1, + ), + ) + assert batch_options.priority_aware_batching_options is not None + self.assertEqual( + batch_options.priority_aware_batching_options.max_queue_depth, -1 + ) + if __name__ == "__main__": absltest.main()