Skip to content

Commit f5bd624

Browse files
samos123Orbax Authors
authored andcommitted
Reduce default GCS target size to 400MiB for OCDBT in GCS
The default target data file size for GCS OCDBT checkpoints is reduced from 2 GiB to 400 MiB for improved write and read performance. This also reduces the impact of a single GCS write taking too long. PiperOrigin-RevId: 908509104
1 parent 218da6d commit f5bd624

4 files changed

Lines changed: 363 additions & 19 deletions

File tree

checkpoint/CHANGELOG.md

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,10 @@ merging.
2020
leaves can also be saved as individual checkpointables.
2121
- Move MTC files to multi_tier_checkpointing and use local checkpoint engine
2222

23+
### Changed
24+
25+
- Reduced default OCDBT target file size from 2GiB to 400MiB for GCS paths.
26+
2327
## [0.11.36] - 2026-04-14
2428

2529
### Added

checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py

Lines changed: 42 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -34,10 +34,10 @@
3434
from orbax.checkpoint._src.metadata import sharding as sharding_metadata
3535
from orbax.checkpoint._src.metadata import value as value_metadata
3636
from orbax.checkpoint._src.path import async_path
37+
from orbax.checkpoint._src.path import gcs_utils
3738
from orbax.checkpoint._src.serialization import types
3839
import tensorstore as ts
3940

40-
4141
JsonSpec: TypeAlias = dict[str, Any]
4242
Shape: TypeAlias = arrays_types.Shape
4343
DType: TypeAlias = arrays_types.DType
@@ -51,6 +51,7 @@
5151
REPLICA_SUBDIR_SUFFIX = 'replica_'
5252
_OCDBT_PROCESS_ID_RE = r'[A-Za-z0-9]+'
5353
_DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE = 2**31 # 2 GiB
54+
_GCS_OCDBT_TARGET_DATA_FILE_SIZE = 400 * 2**20 # 400 MiB
5455

5556
ZARR_VER2 = 'zarr'
5657
ZARR_VER3 = 'zarr3'
@@ -240,19 +241,45 @@ def build_kvstore_tspec_for_merge(
240241
)
241242

242243

244+
def _get_backend_ocdbt_target_data_file_size(
245+
kvstore_spec: JsonSpec | None,
246+
) -> int:
247+
"""Gets OCDBT target data file size based on kvstore spec."""
248+
if kvstore_spec is None:
249+
return _DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE
250+
base = kvstore_spec.get('base')
251+
252+
if isinstance(base, str):
253+
# OCDBT base is generally a string when it's a GCS path.
254+
if gcs_utils.is_gcs_path(epath.Path(base)):
255+
return _GCS_OCDBT_TARGET_DATA_FILE_SIZE
256+
elif isinstance(base, dict):
257+
# OCDBT base can also be a dict with 'driver' and 'path' keys.
258+
if base.get('driver') == 'gcs':
259+
return _GCS_OCDBT_TARGET_DATA_FILE_SIZE
260+
path_str = base.get('path')
261+
if path_str and gcs_utils.is_gcs_path(epath.Path(path_str)):
262+
return _GCS_OCDBT_TARGET_DATA_FILE_SIZE
263+
264+
return _DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE
265+
266+
243267
def add_ocdbt_write_options(
244268
kvstore_tspec: JsonSpec,
245269
target_data_file_size: int | None = None,
246270
) -> None:
247271
"""Adds write-specific options to a TensorStore OCDBT KVStore spec."""
248-
if target_data_file_size is not None:
249-
# TODO: b/354139177 - disallow too small values, too.
250-
if target_data_file_size < 0:
251-
raise ValueError(
252-
'OCDBT target_data_file_size must be >= 0, where 0 means no limit'
253-
f'; got {target_data_file_size}'
254-
)
255-
kvstore_tspec['target_data_file_size'] = target_data_file_size
272+
if target_data_file_size is None:
273+
target_data_file_size = _get_backend_ocdbt_target_data_file_size(
274+
kvstore_tspec
275+
)
276+
# TODO: b/354139177 - Disallow too small values, too.
277+
if target_data_file_size < 0:
278+
raise ValueError(
279+
'OCDBT target_data_file_size must be >= 0, where 0 means no limit'
280+
f'; got {target_data_file_size}'
281+
)
282+
kvstore_tspec['target_data_file_size'] = target_data_file_size
256283

257284
kvstore_tspec['config'] = {
258285
# Store .zarray metadata inline but not large chunks.
@@ -342,13 +369,14 @@ def calculate_chunk_byte_size(
342369
*,
343370
chunk_byte_size: int | None,
344371
ocdbt_target_data_file_size: int | None = None,
372+
kvstore_spec: JsonSpec | None = None,
345373
) -> int | None:
346374
"""Selects chunk byte size to fit both target data file and chunk sizes."""
347375
# Check if the chunk size would exceed ocdbt target file size.
348376
if ocdbt_target_data_file_size is None:
349-
# Set to default used by TensorStore
350-
# (from https://google.github.io/tensorstore/kvstore/ocdbt/index.html).
351-
ocdbt_target_data_file_size = _DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE
377+
ocdbt_target_data_file_size = _get_backend_ocdbt_target_data_file_size(
378+
kvstore_spec
379+
)
352380

353381
if ocdbt_target_data_file_size == 0:
354382
# No limit.
@@ -509,6 +537,7 @@ def __init__(
509537
target_storage_dtype,
510538
chunk_byte_size=chunk_byte_size,
511539
ocdbt_target_data_file_size=ocdbt_target_data_file_size,
540+
kvstore_spec=tspec['kvstore'],
512541
)
513542
# Choose chunk shape.
514543
chunk_shape = subchunking.choose_chunk_shape(
@@ -719,6 +748,7 @@ def get_json_tspec_write(
719748
dtype,
720749
chunk_byte_size=chunk_byte_size,
721750
ocdbt_target_data_file_size=ocdbt_target_data_file_size,
751+
kvstore_spec=tspec['kvstore'],
722752
)
723753

724754
chunk_shape = subchunking.choose_chunk_shape(

checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils_test.py

Lines changed: 147 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -33,12 +33,20 @@ class AddOcdbtWriteOptionsTest(parameterized.TestCase):
3333
def test_ocdbt_target_data_file_size_none(self):
3434
kvstore_tspec = {}
3535
ts_utils.add_ocdbt_write_options(kvstore_tspec, target_data_file_size=None)
36-
self.assertNotIn('target_data_file_size', kvstore_tspec)
36+
self.assertIn('target_data_file_size', kvstore_tspec)
37+
self.assertEqual(
38+
kvstore_tspec['target_data_file_size'],
39+
ts_utils._DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE,
40+
)
3741

3842
def test_ocdbt_target_data_file_size_none_is_default(self):
3943
kvstore_tspec = {}
4044
ts_utils.add_ocdbt_write_options(kvstore_tspec)
41-
self.assertNotIn('target_data_file_size', kvstore_tspec)
45+
self.assertIn('target_data_file_size', kvstore_tspec)
46+
self.assertEqual(
47+
kvstore_tspec['target_data_file_size'],
48+
ts_utils._DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE,
49+
)
4250

4351
def test_ocdbt_target_data_file_size_rejects_negative_value(self):
4452
with self.assertRaises(ValueError):
@@ -58,6 +66,56 @@ def test_ocdbt_target_data_file_size_sets_value(
5866
)
5967

6068

69+
class GetBackendOcdbtTargetDataFileSizeTest(parameterized.TestCase):
70+
71+
@parameterized.named_parameters(
72+
dict(
73+
testcase_name='none_spec',
74+
kvstore_spec=None,
75+
expected_size=ts_utils._DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE,
76+
),
77+
dict(
78+
testcase_name='base_is_gcs_str',
79+
kvstore_spec={'base': 'gs://bucket/path'},
80+
expected_size=ts_utils._GCS_OCDBT_TARGET_DATA_FILE_SIZE,
81+
),
82+
dict(
83+
testcase_name='base_is_local_str',
84+
kvstore_spec={'base': '/tmp/path'},
85+
expected_size=ts_utils._DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE,
86+
),
87+
dict(
88+
testcase_name='base_is_gcs_driver_dict',
89+
kvstore_spec={'base': {'driver': 'gcs', 'bucket': 'bucket'}},
90+
expected_size=ts_utils._GCS_OCDBT_TARGET_DATA_FILE_SIZE,
91+
),
92+
dict(
93+
testcase_name='base_is_dict_with_gcs_path',
94+
kvstore_spec={'base': {'driver': 'file', 'path': 'gs://bucket/path'}},
95+
expected_size=ts_utils._GCS_OCDBT_TARGET_DATA_FILE_SIZE,
96+
),
97+
dict(
98+
testcase_name='base_is_dict_with_local_path',
99+
kvstore_spec={'base': {'driver': 'file', 'path': '/tmp/path'}},
100+
expected_size=ts_utils._DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE,
101+
),
102+
dict(
103+
testcase_name='spec_without_base',
104+
kvstore_spec={},
105+
expected_size=ts_utils._DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE,
106+
),
107+
)
108+
def test_get_backend_ocdbt_target_data_file_size(
109+
self,
110+
kvstore_spec: ts_utils.JsonSpec | None,
111+
expected_size: int,
112+
):
113+
self.assertEqual(
114+
ts_utils._get_backend_ocdbt_target_data_file_size(kvstore_spec),
115+
expected_size,
116+
)
117+
118+
61119
class BuildArrayTSpecForWriteTest(parameterized.TestCase):
62120

63121
def setUp(self):
@@ -231,11 +289,41 @@ def test_ocdbt_kvstore_with_gcs_path(
231289
os.path.join(expected_directory or directory, 'ocdbt.process_0'),
232290
)
233291
self.assertEqual(kvstore_tspec['path'], self.param_name)
292+
self.assertEqual(
293+
kvstore_tspec['target_data_file_size'],
294+
ts_utils._GCS_OCDBT_TARGET_DATA_FILE_SIZE,
295+
)
234296

235-
@parameterized.product(use_zarr3=(True, False))
236-
def test_ocdbt_kvstore_default_target_data_file_size(self, use_zarr3: bool):
297+
@parameterized.named_parameters(
298+
dict(
299+
testcase_name='zarr3_None',
300+
use_zarr3=True,
301+
directory_override=None,
302+
),
303+
dict(
304+
testcase_name='zarr3_tfhub',
305+
use_zarr3=True,
306+
directory_override='/tfhub/prod/model',
307+
),
308+
dict(
309+
testcase_name='no_zarr3_None',
310+
use_zarr3=False,
311+
directory_override=None,
312+
),
313+
dict(
314+
testcase_name='no_zarr3_tfhub',
315+
use_zarr3=False,
316+
directory_override='/tfhub/prod/model',
317+
),
318+
)
319+
def test_ocdbt_kvstore_default_target_data_file_size(
320+
self,
321+
use_zarr3: bool,
322+
directory_override: str | None,
323+
):
324+
directory = directory_override or self.directory
237325
tspec = self.array_write_spec_constructor(
238-
directory=self.directory,
326+
directory=directory,
239327
relative_array_filename=self.param_name,
240328
use_zarr3=use_zarr3,
241329
use_ocdbt=True,
@@ -244,7 +332,11 @@ def test_ocdbt_kvstore_default_target_data_file_size(self, use_zarr3: bool):
244332
self.assertEqual(tspec.metadata.use_zarr3, use_zarr3)
245333
self.assertTrue(tspec.metadata.use_ocdbt)
246334
self.assertEqual(tspec.json['kvstore']['driver'], 'ocdbt')
247-
self.assertNotIn('target_data_file_size', tspec.json['kvstore'])
335+
self.assertIn('target_data_file_size', tspec.json['kvstore'])
336+
self.assertEqual(
337+
tspec.json['kvstore']['target_data_file_size'],
338+
ts_utils._DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE,
339+
)
248340

249341
@parameterized.named_parameters(
250342
dict(testcase_name='none', target_data_file_size=None),
@@ -267,7 +359,11 @@ def test_ocdbt_kvstore_target_data_file_size(
267359
kvstore_tspec = tspec.json['kvstore']
268360
self.assertEqual(kvstore_tspec['driver'], 'ocdbt')
269361
if target_data_file_size is None:
270-
self.assertNotIn('target_data_file_size', kvstore_tspec)
362+
self.assertIn('target_data_file_size', kvstore_tspec)
363+
self.assertEqual(
364+
kvstore_tspec['target_data_file_size'],
365+
ts_utils._DEFAULT_OCDBT_TARGET_DATA_FILE_SIZE,
366+
)
271367
else:
272368
self.assertEqual(
273369
kvstore_tspec['target_data_file_size'], target_data_file_size
@@ -593,6 +689,50 @@ def test_chunk_byte_size_is_adjusted_for_target_data_file_size(
593689
expected_chunk_byte_size_limit,
594690
)
595691

692+
@parameterized.product(
693+
chunk_byte_size=[None, 512 * 2**20],
694+
use_zarr3=[True, False],
695+
)
696+
def test_gcs_chunk_byte_size_is_adjusted_for_target_data_file_size(
697+
self,
698+
chunk_byte_size: int | None,
699+
use_zarr3: bool,
700+
):
701+
gcs_dir = 'gs://gcs_bucket/object_path'
702+
self.shape = (8 * 1024, 2 * 1024, 4 * 1024)
703+
self.write_shape = (2 * 1024, 1024, 2 * 1024)
704+
storage_dtype = self.dtype
705+
706+
tspec = ts_utils.ArrayWriteSpec(
707+
directory=gcs_dir,
708+
relative_array_filename=self.param_name,
709+
global_shape=self.shape,
710+
write_shape=self.write_shape,
711+
dtype=self.dtype,
712+
target_dtype=None,
713+
chunk_byte_size=chunk_byte_size,
714+
use_zarr3=use_zarr3,
715+
use_ocdbt=True,
716+
process_id='w13',
717+
ocdbt_target_data_file_size=None,
718+
)
719+
self.assertEqual(tspec.metadata.dtype, storage_dtype)
720+
chunk_shape = self._get_chunk_shape_from_tspec(
721+
tspec.json,
722+
use_zarr3=use_zarr3,
723+
)
724+
np.testing.assert_array_equal(chunk_shape, tspec.metadata.chunk_shape)
725+
self.assertTrue(
726+
subchunking.validate_divisible_shapes(self.write_shape, chunk_shape),
727+
f'Write shape {self.write_shape} is not divisible by chunk shape'
728+
f' {chunk_shape}.',
729+
)
730+
731+
self.assertLessEqual(
732+
math.prod(chunk_shape) * storage_dtype.itemsize,
733+
ts_utils._GCS_OCDBT_TARGET_DATA_FILE_SIZE,
734+
)
735+
596736
def test_maybe_cloud_storage(self):
597737
gs_path = 'gs://some-buck/path'
598738
gs_spec = serialization.get_tensorstore_spec(gs_path, ocdbt=True)

0 commit comments

Comments
 (0)