Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 33 additions & 0 deletions export/orbax/export/obm_configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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:
Expand Down
62 changes: 62 additions & 0 deletions export/orbax/export/obm_configs_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading