From 332d1a84efe96f81cb25a38baeac6bdaa5513132 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Sat, 19 Sep 2026 12:14:05 +0800 Subject: [PATCH 1/8] [bugfix] slice rows per rank and add data_config.min_batch_size Training crashed with "Expected more than 1 value per channel" from a BatchNorm1d in train mode when a table's row count left every dataloader worker with a single leftover row. calc_slice_position split a table into num_workers * world_size slices and its drop_redundant_bs_eq_one guard only covered the residue where the base slice is a whole number of batches; when floor(rows / slices) % batch_size == 1 the guard did nothing and each worker emitted a 1-row tail batch. The slicing now works per rank: every rank takes the same number of rows of each table or session, whole batches are spread over the rank's workers and the partial tail goes to one worker, so a pass yields one partial batch per rank instead of one per worker. Ranks stay in lockstep by construction and extra rows are only skipped when they would buy a rank an additional step, which removes the residue case analysis and the pre_total_remain chaining. Models that need at least N rows per batch set data_config.min_batch_size, which drops a shorter final batch on every rank identically; it defaults to 0 and replaces drop_remainder's mechanism, which now equals min_batch_size=batch_size. Both apply to train and eval only, so predict reads every row. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01163rPXLGHiM74LxdbWjMox --- docs/source/feature/data.md | 7 + tzrec/datasets/csv_dataset.py | 4 +- tzrec/datasets/dataset.py | 23 +- tzrec/datasets/odps_dataset.py | 98 ++++---- tzrec/datasets/odps_dataset_test.py | 1 + tzrec/datasets/parquet_dataset.py | 28 +-- tzrec/datasets/parquet_dataset_test.py | 131 ++++++++++- tzrec/datasets/utils.py | 304 +++++++++++++++---------- tzrec/datasets/utils_test.py | 187 +++++++++++---- tzrec/protos/data.proto | 4 + 10 files changed, 554 insertions(+), 233 deletions(-) diff --git a/docs/source/feature/data.md b/docs/source/feature/data.md index 74b1505eb..51ff801ca 100644 --- a/docs/source/feature/data.md +++ b/docs/source/feature/data.md @@ -330,6 +330,13 @@ pipeline.global-job-parameters: | ### drop_remainder - 是否丢弃掉最后一个不足batch_size的batch数据,默认为false +- 仅在训练和评估时生效,预测时不会丢弃任何数据;等价于`min_batch_size`设置为`batch_size` + +### min_batch_size + +- 训练和评估时,丢弃掉每个数据读取进程最后一个行数小于`min_batch_size`的batch,默认为0(不丢弃) +- 使用BatchNorm等要求batch内至少2行样本的模型时,建议设置为2,避免样本表行数恰好使最后一个batch只剩1行导致训练失败 +- 注:OdpsDataset和ParquetDataset按行切分数据时,会保证每个`proc`(rank)读取相同的步数,以避免同步训练时卡住;为此每张表(或每个分区)最多有`nproc - 1`行样本不会被读取,预测时不受影响 ### batch_cost_size diff --git a/tzrec/datasets/csv_dataset.py b/tzrec/datasets/csv_dataset.py index 4af8fe326..917f4c165 100644 --- a/tzrec/datasets/csv_dataset.py +++ b/tzrec/datasets/csv_dataset.py @@ -68,9 +68,10 @@ def __init__( input_path, self._batch_size, list(self._selected_input_names) if self._selected_input_names else None, - self._data_config.drop_remainder, + self._drop_remainder, shuffle=self._data_config.shuffle and self._mode == Mode.TRAIN, shuffle_buffer_size=self._data_config.shuffle_buffer_size, + min_batch_size=self._min_batch_size, column_names=column_names, delimiter=self._data_config.delimiter, column_types=column_types, @@ -119,6 +120,7 @@ def __init__( shuffle_buffer_size, sample_cost_field=sample_cost_field, batch_cost_size=batch_cost_size, + **kwargs, ) self._csv_fmt = ds.CsvFileFormat( parse_options=pa.csv.ParseOptions(delimiter=delimiter), diff --git a/tzrec/datasets/dataset.py b/tzrec/datasets/dataset.py index 11972bbad..c7229c4bc 100644 --- a/tzrec/datasets/dataset.py +++ b/tzrec/datasets/dataset.py @@ -208,6 +208,18 @@ def __init__( self._batch_size = config_util.get_inference_batch_size(data_config) else: self._batch_size = data_config.batch_size + # predict keeps every row: no tail dropping and no rank equalization + if mode == Mode.PREDICT: + self._drop_remainder = False + self._min_batch_size = 0 + else: + self._drop_remainder = data_config.drop_remainder + self._min_batch_size = data_config.min_batch_size + if self._min_batch_size > self._batch_size: + raise ValueError( + f"data_config.min_batch_size[{self._min_batch_size}] must not " + f"exceed the batch size[{self._batch_size}]." + ) self._sampler = None self._sampler_inited = False @@ -533,11 +545,14 @@ class BaseReader(metaclass=_reader_meta_cls): input_path (str): data input path. batch_size (int): batch size. selected_cols (list): selection column names. - drop_remainder (bool): drop last batch. + drop_remainder (bool): drop last batch, same as min_batch_size=batch_size. shuffle (bool): shuffle data or not. shuffle_buffer_size (int): buffer size for shuffle. sample_cost_field (str): sample cost field name. batch_cost_size (int): batch cost limit size. + min_batch_size (int): drop a final batch with fewer rows, 0 disables. + equalize_rank_steps (bool): make every rank yield the same batch sizes, + honored by readers that slice rows by count. """ def __init__( @@ -550,12 +565,16 @@ def __init__( shuffle_buffer_size: int = 32, sample_cost_field: Optional[str] = None, batch_cost_size: Optional[int] = None, + min_batch_size: int = 0, + equalize_rank_steps: bool = False, **kwargs: Any, ) -> None: self._input_path = input_path self._batch_size = batch_size self._selected_cols = selected_cols self._drop_remainder = drop_remainder + self._min_batch_size = batch_size if drop_remainder else min_batch_size + self._equalize_rank_steps = equalize_rank_steps self._shuffle = shuffle self._shuffle_buffer_size = shuffle_buffer_size self._sample_cost_field = sample_cost_field @@ -624,7 +643,7 @@ def _arrow_reader_iter( [buff_data, pa.Table.from_batches([read_data])] ) except StopIteration: - if self._drop_remainder or buff_data is None: + if buff_data is None or len(buff_data) < self._min_batch_size: data = buff_data = None else: data, buff_data = self._slice_buff_data(buff_data) diff --git a/tzrec/datasets/odps_dataset.py b/tzrec/datasets/odps_dataset.py index fe594fc07..d992f1223 100644 --- a/tzrec/datasets/odps_dataset.py +++ b/tzrec/datasets/odps_dataset.py @@ -387,13 +387,14 @@ def __init__( input_path, self._batch_size, list(self._selected_input_names) if self._selected_input_names else None, - self._data_config.drop_remainder, + self._drop_remainder, is_orderby_partition=self._data_config.is_orderby_partition, quota_name=self._data_config.odps_data_quota_name, - drop_redundant_bs_eq_one=self._mode != Mode.PREDICT, compression=self._data_config.odps_data_compression, sample_cost_field=self._data_config.sample_cost_field, batch_cost_size=self._data_config.batch_cost_size, + min_batch_size=self._min_batch_size, + equalize_rank_steps=self._mode != Mode.PREDICT, ) @@ -404,16 +405,16 @@ class OdpsReader(BaseReader): input_path (str): data input path. batch_size (int): batch size. selected_cols (list): selection column names. - drop_remainder (bool): drop last batch less than batch_size. + drop_remainder (bool): drop last batch, same as min_batch_size=batch_size. shuffle (bool): shuffle data or not. shuffle_buffer_size (int): buffer size for shuffle. is_orderby_partition (bool): read data order by table partitions or not. quota_name (str): storage api quota name. - drop_redundant_bs_eq_one (bool): drop last redundant batch with batch_size - equal one to prevent train_eval hung. compression (str): storage api data compression name. sample_cost_field (str): sample cost field name. batch_cost_size (int): batch cost limit size. + min_batch_size (int): drop a final batch with fewer rows, 0 disables. + equalize_rank_steps (bool): make every rank yield the same batch sizes. """ def __init__( @@ -426,7 +427,6 @@ def __init__( shuffle_buffer_size: int = 32, is_orderby_partition: bool = False, quota_name: str = "pay-as-you-go", - drop_redundant_bs_eq_one: bool = False, compression: str = "LZ4_FRAME", sample_cost_field: Optional[str] = None, batch_cost_size: Optional[int] = None, @@ -441,13 +441,13 @@ def __init__( shuffle_buffer_size, sample_cost_field=sample_cost_field, batch_cost_size=batch_cost_size, + **kwargs, ) self._pg = dist_util.get_dist_object_pg() self._is_orderby_partition = is_orderby_partition self._quota_name = quota_name self._compression = _get_compression_type(compression) os.environ["STORAGE_API_QUOTA_NAME"] = quota_name - self._drop_redundant_bs_eq_one = drop_redundant_bs_eq_one self._account, self._odps_endpoint = _create_odps_account() self._proj_to_o = {} @@ -620,61 +620,45 @@ def to_batches( self, worker_id: int = 0, num_workers: int = 1 ) -> Iterator[Dict[str, pa.Array]]: """Get batch iterator.""" - input_paths = self._input_path.split(",") - num_tables = len(input_paths) - - def _combined_reader() -> Iterator[pa.RecordBatch]: - remain_row_count = 0 - - for table_idx, input_path in enumerate(input_paths): - is_last_table = table_idx == num_tables - 1 - _, table_name, _, _ = _parse_table_path(input_path) - client = self._table_to_cli[table_name] - sess_reqs = self._input_to_sess[input_path] - num_sess = len(sess_reqs) - - for sess_idx, sess_req in enumerate(sess_reqs): - is_last_session = sess_idx == num_sess - 1 - # Only drop redundant on the very last session of the very - # last table - should_drop_redundant = ( - self._drop_redundant_bs_eq_one - and is_last_table - and is_last_session + # (source_id_prefix, record_count, client, session) in read order + sources = [] + for input_path in self._input_path.split(","): + _, table_name, _, _ = _parse_table_path(input_path) + client = self._table_to_cli[table_name] + for sess_req in self._input_to_sess[input_path]: + sources.append( + ( + f"{input_path}#{sess_req.session_id}", + _get_session_record_count(client, sess_req), + client, + sess_req, ) + ) + intervals = calc_slice_intervals( + [(prefix, record_count) for prefix, record_count, _, _ in sources], + worker_id, + num_workers, + self._batch_size, + self._equalize_rank_steps, + self._min_batch_size, + checkpoint_state=self._checkpoint_state, + ) - # Get session record count - record_count = _get_session_record_count(client, sess_req) - - # Generate source_id with session_id for unique identification - source_id_prefix = f"{input_path}#{sess_req.session_id}" - - # Calculate intervals (similar to parquet pattern) - worker_intervals, remain_row_count = calc_slice_intervals( - record_count, - worker_id, - num_workers, + def _combined_reader() -> Iterator[pa.RecordBatch]: + for prefix, _, client, sess_req in sources: + for start, end in intervals[prefix]: + if start >= end: + continue + yield from _reader_iter( + client, + sess_req, self._batch_size, - should_drop_redundant, - pre_total_remain=remain_row_count, - checkpoint_state=self._checkpoint_state, - input_path=source_id_prefix, + self._compression, + start, + end, + f"{prefix}:{start}", ) - for start, end in worker_intervals: - if start >= end: - continue - source_id = f"{source_id_prefix}:{start}" - yield from _reader_iter( - client, - sess_req, - self._batch_size, - self._compression, - start, - end, - source_id, - ) - yield from self._arrow_reader_iter(_combined_reader()) diff --git a/tzrec/datasets/odps_dataset_test.py b/tzrec/datasets/odps_dataset_test.py index 461b00861..2ce50dca0 100644 --- a/tzrec/datasets/odps_dataset_test.py +++ b/tzrec/datasets/odps_dataset_test.py @@ -387,6 +387,7 @@ def test_odps_dataset_checkpoint_resume_orderby_partition(self): label_fields=["label"], is_orderby_partition=True, odps_data_quota_name=self.test_quota, + min_batch_size=2, ), features=features, input_path=input_path, diff --git a/tzrec/datasets/parquet_dataset.py b/tzrec/datasets/parquet_dataset.py index 69bbdc3eb..0d48712e9 100644 --- a/tzrec/datasets/parquet_dataset.py +++ b/tzrec/datasets/parquet_dataset.py @@ -137,12 +137,13 @@ def __init__( input_path, self._batch_size, list(self._selected_input_names) if self._selected_input_names else None, - self._data_config.drop_remainder, + self._drop_remainder, shuffle=self._data_config.shuffle and self._mode == Mode.TRAIN, shuffle_buffer_size=self._data_config.shuffle_buffer_size, - drop_redundant_bs_eq_one=self._mode != Mode.PREDICT, sample_cost_field=self._data_config.sample_cost_field, batch_cost_size=self._data_config.batch_cost_size, + min_batch_size=self._min_batch_size, + equalize_rank_steps=self._mode != Mode.PREDICT, ) @@ -153,14 +154,14 @@ class ParquetReader(BaseReader): input_path (str): data input path. batch_size (int): batch size. selected_cols (list): selection column names. - drop_remainder (bool): drop last batch. + drop_remainder (bool): drop last batch, same as min_batch_size=batch_size. shuffle (bool): shuffle data or not. shuffle_buffer_size (int): buffer size for shuffle. - drop_redundant_bs_eq_one (bool): drop last redundant batch with batch_size - equal one to prevent train_eval hung. - rebalance (bool): rebalance parquet rows to equal number for each worker. + rebalance (bool): rebalance parquet rows to equal number for each rank. sample_cost_field (str): sample cost field name. batch_cost_size (int): batch cost limit size. + min_batch_size (int): drop a final batch with fewer rows, 0 disables. + equalize_rank_steps (bool): make every rank yield the same batch sizes. """ def __init__( @@ -171,7 +172,6 @@ def __init__( drop_remainder: bool = False, shuffle: bool = False, shuffle_buffer_size: int = 32, - drop_redundant_bs_eq_one: bool = False, rebalance: bool = True, sample_cost_field: Optional[str] = None, batch_cost_size: Optional[int] = None, @@ -186,9 +186,9 @@ def __init__( shuffle_buffer_size, sample_cost_field=sample_cost_field, batch_cost_size=batch_cost_size, + **kwargs, ) self._pg = dist_util.get_dist_object_pg() - self._drop_redundant_bs_eq_one = drop_redundant_bs_eq_one self._rebalance = rebalance self._ordered_cols = None @@ -248,18 +248,20 @@ def to_batches( worker_intervals = [(0, sys.maxsize)] if self._rebalance: - worker_intervals, _ = calc_slice_intervals( - sum(self._num_rows), + worker_intervals = calc_slice_intervals( + [(self._input_path, sum(self._num_rows))], worker_id, num_workers, self._batch_size, - self._drop_redundant_bs_eq_one, + self._equalize_rank_steps, + self._min_batch_size, checkpoint_state=self._checkpoint_state, - input_path=self._input_path, - ) + )[self._input_path] def _combined_reader() -> Iterator[pa.RecordBatch]: for start, end in worker_intervals: + if start >= end: + continue yield from _reader_iter( self._input_files if self._rebalance diff --git a/tzrec/datasets/parquet_dataset_test.py b/tzrec/datasets/parquet_dataset_test.py index 64b75620d..370bdc870 100644 --- a/tzrec/datasets/parquet_dataset_test.py +++ b/tzrec/datasets/parquet_dataset_test.py @@ -23,13 +23,14 @@ from torch import distributed as dist from torch.utils.data import DataLoader +from tzrec.constant import Mode from tzrec.datasets.dataset import create_dataloader from tzrec.datasets.parquet_dataset import ParquetDataset, ParquetReader, ParquetWriter from tzrec.features.feature import create_features from tzrec.protos import data_pb2, feature_pb2 from tzrec.utils import misc_util from tzrec.utils.checkpoint_util import EPOCHS_COMPLETED, update_dataloder_state -from tzrec.utils.test_util import make_test_dir +from tzrec.utils.test_util import make_test_dir, parameterized_name_func class ParquetDatasetTest(unittest.TestCase): @@ -312,6 +313,78 @@ def _drain(iterator): self.assertEqual(num_meta_only_rows, num_total_rows) del dataloader3 + @parameterized.expand( + [ + # 8200 rows over 8 workers: one 8-row tail instead of eight 1-row tails + [0, [1024] * 8 + [8]], + [2, [1024] * 8 + [8]], + # drop_remainder drops the 8-row tail only + [1024, [1024] * 8], + ], + name_func=parameterized_name_func, + ) + def test_create_dataloader_tail_batches(self, min_batch_size, expected): + feature_cfgs = self._create_feature_cfgs() + features = create_features(feature_cfgs) + with tempfile.TemporaryDirectory(prefix="tzrec_") as test_dir: + self._create_test_parquet_data(test_dir, num_rows=8200) + data_config = data_pb2.DataConfig( + batch_size=1024, + dataset_type=data_pb2.DatasetType.ParquetDataset, + fg_mode=data_pb2.FgMode.FG_NONE, + label_fields=["label"], + num_workers=8, + min_batch_size=min_batch_size, + drop_remainder=min_batch_size == 1024, + ) + dataloader = create_dataloader(data_config, features, f"{test_dir}/*") + sizes = sorted( + (len(batch.labels["label"]) for batch in dataloader.get_iterator()), + reverse=True, + ) + self.assertEqual(sizes, expected) + + def test_create_dataloader_predict_keeps_every_row(self): + feature_cfgs = self._create_feature_cfgs() + features = create_features(feature_cfgs) + with tempfile.TemporaryDirectory(prefix="tzrec_") as test_dir: + self._create_test_parquet_data(test_dir, num_rows=8201) + data_config = data_pb2.DataConfig( + batch_size=1024, + dataset_type=data_pb2.DatasetType.ParquetDataset, + fg_mode=data_pb2.FgMode.FG_NONE, + label_fields=["label"], + num_workers=8, + drop_remainder=True, + min_batch_size=2, + ) + dataloader = create_dataloader( + data_config, + features, + f"{test_dir}/*", + reserved_columns=["label"], + mode=Mode.PREDICT, + ) + num_rows = sum( + len(batch.reserves.get()) for batch in dataloader.get_iterator() + ) + self.assertEqual(num_rows, 8201) + + def test_min_batch_size_exceeds_batch_size(self): + feature_cfgs = self._create_feature_cfgs() + features = create_features(feature_cfgs) + with tempfile.TemporaryDirectory(prefix="tzrec_") as test_dir: + self._create_test_parquet_data(test_dir, num_rows=10) + data_config = data_pb2.DataConfig( + batch_size=4, + dataset_type=data_pb2.DatasetType.ParquetDataset, + fg_mode=data_pb2.FgMode.FG_NONE, + label_fields=["label"], + min_batch_size=5, + ) + with self.assertRaisesRegex(ValueError, "min_batch_size"): + ParquetDataset(data_config, features, f"{test_dir}/*") + def test_create_dataloader_get_iterator_reuses_eager_iterator(self): """get_iterator() yields the eagerly prefetched iterator first. @@ -436,6 +509,62 @@ def _reader_worker(rank, port): if p.exitcode != 0: raise RuntimeError(f"reader worker-{i} failed.") + @parameterized.expand( + [ + # extra row would buy rank 0 a fifth step: dropped + [33, 4, 0, [4, 4], [[4] * 4, [4] * 4]], + # residue-1 tails stay when min_batch_size is 0 + [35, 4, 0, [5, 5], [[4] * 4 + [2], [4] * 4 + [1]]], + # ... and go together with the extra row when min_batch_size=2 + [35, 4, 2, [4, 4], [[4] * 4, [4] * 4]], + # 8200 rows over two ranks: 4100 each, 4-row tails + [8200, 1024, 2, [5, 5], [[1024] * 4 + [4], [1024] * 4 + [4]]], + ], + name_func=parameterized_name_func, + ) + def test_parquet_reader_equal_steps( + self, num_rows, batch_size, min_batch_size, steps, sizes + ): + def _reader_worker(rank, port, queue): + os.environ["RANK"] = str(rank) + os.environ["WORLD_SIZE"] = str(2) + os.environ["MASTER_ADDR"] = "127.0.0.1" + os.environ["MASTER_PORT"] = str(port) + dist.init_process_group(backend="gloo") + reader = ParquetReader( + os.path.join(self.test_dir, "*.parquet"), + batch_size=batch_size, + min_batch_size=min_batch_size, + equalize_rank_steps=True, + ) + queue.put((rank, [len(b["id_a"]) for b in reader.to_batches(rank, 2)])) + + t = pa.Table.from_arrays( + [pa.array(["1"] * num_rows), pa.array([0] * num_rows)], + names=["id_a", "label"], + ) + writer = parquet.ParquetWriter( + os.path.join(self.test_dir, "part-0.parquet"), schema=t.schema + ) + writer.write_table(t) + writer.close() + + port = misc_util.get_free_port() + queue = mp.Queue() + procs = [ + mp.Process(target=_reader_worker, args=(rank, port, queue)) + for rank in range(2) + ] + for p in procs: + p.start() + results = dict(queue.get() for _ in procs) + for i, p in enumerate(procs): + p.join() + if p.exitcode != 0: + raise RuntimeError(f"reader worker-{i} failed.") + self.assertEqual([len(results[r]) for r in range(2)], steps) + self.assertEqual([results[r] for r in range(2)], sizes) + class ParquetWriterTest(unittest.TestCase): def setUp(self): diff --git a/tzrec/datasets/utils.py b/tzrec/datasets/utils.py index 6dc5cdeff..b89a67d21 100644 --- a/tzrec/datasets/utils.py +++ b/tzrec/datasets/utils.py @@ -10,6 +10,7 @@ # limitations under the License. import glob +import os import re from dataclasses import dataclass, field from typing import Dict, List, Optional, Tuple @@ -759,65 +760,6 @@ def combine_negs_to_candidate_sequence( return pa.array(rows, type=pa.string()), pos_lengths -def calc_slice_position( - row_count: int, - slice_id: int, - slice_count: int, - batch_size: int, - drop_redundant_bs_eq_one: bool, - pre_total_remain: int = 0, -) -> Tuple[int, int, int]: - """Calc table read position according to the slice information. - - Args: - row_count (int): table total row count. - slice_id (int): worker id. - slice_count (int): total worker number. - batch_size (int): batch_size. - drop_redundant_bs_eq_one (bool): drop last redundant batch with batch_size - equal one to prevent train_eval hung. - pre_total_remain (int): remaining total count in pre-table is - insufficient to meet the batch_size requirement for each worker. - - Return: - start (int): start row position in table. - end (int): start row position in table. - total_remain (int): remaining total count in curr-table is - insufficient to meet the batch_size requirement for each worker. - """ - pre_remain_size = int(pre_total_remain / slice_count) - pre_remain_split_point = pre_total_remain % slice_count - - size = int((row_count + pre_total_remain) / slice_count) - split_point = (row_count + pre_total_remain) % slice_count - if slice_id < split_point: - start = slice_id * (size + 1) - end = start + (size + 1) - else: - start = split_point * (size + 1) + (slice_id - split_point) * size - end = start + size - - real_start = ( - start - pre_remain_size * slice_id - min(pre_remain_split_point, slice_id) - ) - real_end = ( - end - - pre_remain_size * (slice_id + 1) - - min(pre_remain_split_point, slice_id + 1) - ) - # when (end - start) % bz = 1 on some workers and - # (end - start) % bz = 0 on other workers, train_eval will hang - if ( - drop_redundant_bs_eq_one - and split_point != 0 - and (end - start) % batch_size == 1 - and size % batch_size == 0 - ): - real_end = real_end - 1 - split_point = 0 - return real_start, real_end, (size % batch_size) * slice_count + split_point - - def calc_remaining_intervals( checkpoint_state: Optional[Dict[str, int]], input_path: str, @@ -876,79 +818,201 @@ def calc_remaining_intervals( return remaining if remaining else [] +def _get_world_size() -> int: + if dist.is_initialized(): + return dist.get_world_size() + return int(os.environ.get("WORLD_SIZE", 1)) + + +def plan_rank_worker_intervals( + source_rows: List[int], + rank: int, + world_size: int, + worker_id: int, + num_workers: int, + batch_size: int, + equalize_rank_steps: bool = False, + min_batch_size: int = 0, +) -> List[Tuple[int, int]]: + """Plan the row range one dataloader worker reads from every source. + + Sources (tables, sessions, or remaining checkpoint intervals) are consumed in + order by a single buffered reader per worker, so a worker's partial batches + are decided by its cumulative row count over all sources. The plan is + computed identically on every rank: + + 1. Every rank takes ``rows // world_size`` rows of each source, laid out + contiguously; the ``rows % world_size`` extra rows are read one each by + the lowest ranks. + 2. Inside a rank, whole batches are spread over its workers (rotating the + first worker across sources) and the partial tail goes to the worker laid + out last, so a worker never manufactures a partial batch of its own. + 3. A cumulative tail smaller than ``min_batch_size`` is dropped. + 4. With ``equalize_rank_steps`` every rank yields the same batch-size list: + the extra rows landing on a worker are kept only when they cannot push + its last batch past ``batch_size``. + + Dropped rows are never read. + + Args: + source_rows (list): row count of every source, in read order. + rank (int): rank of this process. + world_size (int): number of ranks. + worker_id (int): dataloader worker id within the rank. + num_workers (int): dataloader workers per rank. + batch_size (int): batch size. + equalize_rank_steps (bool): make every rank yield the same batch sizes. + min_batch_size (int): drop a final batch with fewer rows, 0 disables. + + Returns: + (start, end) per source for this rank and worker, ``start == end`` when + the worker reads nothing from that source. + """ + assert 0 < num_workers and 0 <= worker_id < num_workers + assert 0 < world_size and 0 <= rank < world_size + assert 0 <= min_batch_size <= batch_size + + # chunk layout of each source inside a rank: (worker, rows) in layout order + layouts: List[List[Tuple[int, int]]] = [] + last_workers: List[int] = [] + cursor = 0 + for rows in source_rows: + full, tail = divmod(rows // world_size, batch_size) + order = [(cursor + i) % num_workers for i in range(num_workers)] + q, r = divmod(full, num_workers) + layout = [ + (w, (q + 1 if i < r else q) * batch_size) for i, w in enumerate(order) + ] + layout[-1] = (order[-1], layout[-1][1] + tail) + layouts.append(layout) + last_workers.append(order[-1]) + cursor = (cursor + full + (1 if tail else 0)) % num_workers + + totals = [0] * num_workers + for layout in layouts: + for w, cnt in layout: + totals[w] += cnt + + # drop a cumulative tail shorter than min_batch_size from the end of the + # worker's stream, walking its chunks backwards + shrink: Dict[Tuple[int, int], int] = {} + for w in range(num_workers): + remain = totals[w] % batch_size + if 0 < remain < min_batch_size: + totals[w] -= remain + for t in reversed(range(len(source_rows))): + cnt = dict(layouts[t])[w] + take = min(cnt, remain) + if take: + shrink[(t, w)] = take + remain -= take + if remain == 0: + break + + extras = [rows % world_size for rows in source_rows] + if equalize_rank_steps: + landing = [0] * num_workers + for t, e in enumerate(extras): + if e: + landing[last_workers[t]] += 1 + for t, e in enumerate(extras): + w = last_workers[t] + remain = totals[w] % batch_size + if e and (remain == 0 or remain + landing[w] > batch_size): + extras[t] = 0 + + result = [] + for t, rows in enumerate(source_rows): + n = rows // world_size + pos = rank * n + min(rank, extras[t]) + start = end = pos + for i, (w, cnt) in enumerate(layouts[t]): + if i == num_workers - 1 and rank < extras[t]: + cnt += 1 + if w == worker_id: + start = pos + end = pos + cnt - shrink.get((t, w), 0) + pos += cnt + result.append((start, end)) + return result + + +def _map_logical_range( + intervals: List[Tuple[int, int]], logical_start: int, logical_end: int +) -> List[Tuple[int, int]]: + """Map a range over the concatenated intervals back to row intervals.""" + result = [] + current_pos = 0 + for interval_start, interval_end in intervals: + interval_len = interval_end - interval_start + overlap_start = max(logical_start, current_pos) + overlap_end = min(logical_end, current_pos + interval_len) + if overlap_start < overlap_end: + result.append( + ( + interval_start + overlap_start - current_pos, + interval_start + overlap_end - current_pos, + ) + ) + current_pos += interval_len + if current_pos >= logical_end: + break + return result + + def calc_slice_intervals( - total_rows: int, + sources: List[Tuple[str, int]], worker_id: int, num_workers: int, - batch_size: int = 1, - drop_redundant_bs_eq_one: bool = False, - pre_total_remain: int = 0, + batch_size: int, + equalize_rank_steps: bool = False, + min_batch_size: int = 0, checkpoint_state: Optional[Dict[str, int]] = None, - input_path: Optional[str] = None, -) -> Tuple[List[Tuple[int, int]], int]: - """Redistribute remaining intervals among workers. +) -> Dict[str, List[Tuple[int, int]]]: + """Assign the row intervals of every source to one global dataloader worker. - Flattens all intervals into a total row count, then assigns a portion - to each worker based on worker_id and num_workers. + ``worker_id`` and ``num_workers`` are the rank-major global ids from + ``BaseDataset.get_worker_info``. The rows each source still has after the + checkpoint state is applied are planned per rank and worker by + ``plan_rank_worker_intervals`` and mapped back onto the remaining intervals. Args: - total_rows (int): total number of rows in the dataset. - worker_id: Current worker's ID (0-indexed). - num_workers: Total number of workers. - batch_size: batch_size. - drop_redundant_bs_eq_one: drop last redundant batch with batch_size - equal one to prevent train_eval hung. - pre_total_remain (int): remaining total count in pre-table is - insufficient to meet the batch_size requirement for each worker. + sources (list): (source_id_prefix, total_rows) in read order. + worker_id (int): global worker id. + num_workers (int): total worker number over all ranks. + batch_size (int): batch size. + equalize_rank_steps (bool): make every rank yield the same batch sizes. + min_batch_size (int): drop a final batch with fewer rows, 0 disables. checkpoint_state (dict): dict mapping source_id to max consumed row index. - input_path (str): the input path to filter checkpoint entries. Returns: - worker_intervals (list): List of (start, end) tuples assigned to this worker. - total_remain (int): remaining total count in curr-table is - insufficient to meet the batch_size requirement for each worker. + dict mapping source_id_prefix to the (start, end) intervals of this worker. """ - intervals: List[Tuple[int, int]] = [] - if checkpoint_state: - intervals = calc_remaining_intervals(checkpoint_state, input_path, total_rows) - total_rows = sum(end - start for start, end in intervals) - - # Reuse calc_slice_position for worker start/end calculation - worker_start, worker_end, total_remain = calc_slice_position( - row_count=total_rows, - slice_id=worker_id, - slice_count=num_workers, - batch_size=batch_size, - drop_redundant_bs_eq_one=drop_redundant_bs_eq_one, - pre_total_remain=pre_total_remain, + world_size = _get_world_size() + assert num_workers % world_size == 0, ( + f"num_workers[{num_workers}] must be a multiple of world_size[{world_size}]" ) - - if checkpoint_state: - # Map worker's logical range [worker_start, worker_end) to actual intervals - result = [] - current_pos = 0 - for interval_start, interval_end in intervals: - interval_len = interval_end - interval_start - interval_logical_start = current_pos - interval_logical_end = current_pos + interval_len - - # Check if this interval overlaps with worker's range - overlap_start = max(worker_start, interval_logical_start) - overlap_end = min(worker_end, interval_logical_end) - - if overlap_start < overlap_end: - # Map back to actual row indices - actual_start = interval_start + (overlap_start - interval_logical_start) - actual_end = interval_start + (overlap_end - interval_logical_start) - result.append((actual_start, actual_end)) - - current_pos = interval_logical_end - if current_pos >= worker_end: - break - else: - result = [(worker_start, worker_end)] - - return result, total_remain + local_workers = num_workers // world_size + rank, local_worker_id = divmod(worker_id, local_workers) + + remaining = [ + calc_remaining_intervals(checkpoint_state, prefix, total_rows) + for prefix, total_rows in sources + ] + plan = plan_rank_worker_intervals( + [sum(end - start for start, end in intervals) for intervals in remaining], + rank, + world_size, + local_worker_id, + local_workers, + batch_size, + equalize_rank_steps, + min_batch_size, + ) + return { + prefix: _map_logical_range(intervals, start, end) + for (prefix, _), intervals, (start, end) in zip(sources, remaining, plan) + } def remove_nullable(field_type: pa.DataType) -> pa.DataType: diff --git a/tzrec/datasets/utils_test.py b/tzrec/datasets/utils_test.py index 3e95a52a2..ac73c6dc9 100644 --- a/tzrec/datasets/utils_test.py +++ b/tzrec/datasets/utils_test.py @@ -10,7 +10,10 @@ # limitations under the License. +import itertools +import random import unittest +from typing import List, Tuple import numpy as np import pyarrow as pa @@ -21,34 +24,140 @@ build_sampler_input, calc_remaining_intervals, calc_slice_intervals, - calc_slice_position, combine_negs_to_candidate_sequence, get_input_fields_proto, + plan_rank_worker_intervals, ) from tzrec.protos import data_pb2 from tzrec.protos.data_pb2 import FieldType +from tzrec.utils.test_util import parameterized_name_func class DatasetUtilsTest(unittest.TestCase): - def test_calc_slice_position(self): - num_tables = 81 - num_workers = 8 - batch_size = 10 - remain_row_counts = [0] * num_workers - worker_row_counts = [0] * num_workers - for i in range(num_tables): - for j in range(num_workers): - start, end, remain_row_counts[j] = calc_slice_position( - row_count=81, - slice_id=j, - slice_count=num_workers, - batch_size=batch_size, - drop_redundant_bs_eq_one=True if i == num_tables - 1 else False, - pre_total_remain=remain_row_counts[j], + @staticmethod + def _rank_batches( + rows: List[int], + world_size: int, + num_workers: int, + batch_size: int, + equalize: bool, + min_batch_size: int, + ) -> Tuple[List[List[int]], int]: + """Simulate every worker's buffered stream; return per-rank sizes, rows read.""" + per_rank = [] + num_read = 0 + seen = [set() for _ in rows] + for rank in range(world_size): + batches = [] + for worker in range(num_workers): + ranges = plan_rank_worker_intervals( + rows, + rank, + world_size, + worker, + num_workers, + batch_size, + equalize, + min_batch_size, ) - worker_row_counts[j] += end - start - self.assertTrue(np.all(np.ceil(np.array(worker_row_counts) / batch_size) == 82)) - self.assertEqual(sum(worker_row_counts), num_tables * 81 - 1) + total = 0 + for t, (start, end) in enumerate(ranges): + assert 0 <= start <= end <= rows[t] + assert seen[t].isdisjoint(range(start, end)) + seen[t].update(range(start, end)) + total += end - start + num_read += total + batches.extend([batch_size] * (total // batch_size)) + if total % batch_size >= max(min_batch_size, 1): + batches.append(total % batch_size) + per_rank.append(sorted(batches)) + return per_rank, num_read + + def test_plan_rank_worker_intervals_invariants(self): + rng = random.Random(0) + for batch_size, world_size, num_workers, equalize in itertools.product( + (1, 2, 3, 4, 8), (1, 2, 3, 4), (1, 2, 3, 4), (False, True) + ): + for min_batch_size in sorted({0, 1, min(2, batch_size), batch_size}): + cases = [ + [rng.randrange(3 * batch_size * world_size + 5) for _ in range(n)] + for n in (1, 2, 3) + for _ in range(20) + ] + cases += [[r] for r in range(4 * batch_size * world_size + 3)] + for rows in cases: + per_rank, num_read = self._rank_batches( + rows, + world_size, + num_workers, + batch_size, + equalize, + min_batch_size, + ) + msg = ( + f"rows={rows} world={world_size} workers={num_workers} " + f"bs={batch_size} min_bs={min_batch_size}" + ) + for batches in per_rank: + self.assertTrue(all(b <= batch_size for b in batches), msg) + self.assertTrue(all(b >= min_batch_size for b in batches), msg) + if equalize: + self.assertEqual(len({len(b) for b in per_rank}), 1, msg) + if not equalize and min_batch_size == 0: + self.assertEqual(num_read, sum(rows), msg) + else: + max_drop = (world_size - 1) * len(rows) + max_drop += ( + max(min_batch_size - 1, 0) * num_workers * world_size + ) + self.assertLessEqual(sum(rows) - num_read, max_drop, msg) + + @parameterized.expand( + [ + # the failing job: 8200 rows, 8 workers -> one 8-row tail, no drop + [[8200], 1, 8, 1024, 0, 9, [8], 8200], + [[8200], 1, 8, 1024, 2, 9, [8], 8200], + [[8201], 1, 8, 1024, 2, 9, [9], 8201], + # extra row would buy a whole step on rank 0 -> drop it + [[33], 2, 4, 4, 0, 4, [], 32], + # residue-1 tails are kept without min_batch_size + [[34], 2, 4, 4, 0, 5, [1], 34], + [[35], 2, 4, 4, 0, 5, [2], 35], + # ... and dropped together with the extras when min_batch_size=2 + [[35], 2, 4, 4, 2, 4, [], 32], + [[36], 2, 4, 4, 2, 5, [2], 36], + # too few rows for one 2-row batch per rank -> empty pass + [[3], 2, 4, 4, 2, 0, [], 0], + # tails of two sources (1000 + 25) sum to 1024 + 1 on one worker + [[1000, 1049], 1, 1, 1024, 0, 3, [1], 2049], + [[1000, 1049], 1, 1, 1024, 2, 2, [], 2048], + # drop_remainder: no partial batch at all + [[8200], 1, 8, 1024, 1024, 8, [], 8192], + ], + name_func=parameterized_name_func, + ) + def test_plan_rank_worker_intervals( + self, + rows, + world_size, + num_workers, + batch_size, + min_batch_size, + steps, + rank0_tails, + num_read, + ): + per_rank, actual_read = self._rank_batches( + rows, world_size, num_workers, batch_size, True, min_batch_size + ) + self.assertEqual([len(b) for b in per_rank], [steps] * world_size) + self.assertEqual([b for b in per_rank[0] if b != batch_size], rank0_tails) + self.assertEqual(actual_read, num_read) + + def test_plan_rank_worker_intervals_predict_keeps_rows(self): + per_rank, num_read = self._rank_batches([35], 2, 4, 4, False, 0) + self.assertEqual([len(b) for b in per_rank], [5, 5]) + self.assertEqual(num_read, 35) def test_calc_remaining_intervals_no_checkpoint(self): """Test remaining intervals when no checkpoint exists.""" @@ -129,13 +238,13 @@ def test_calc_slice_intervals_single_worker(self): "/data/test.parquet:0": 99, "/data/test.parquet:500": 599, } - result, _ = calc_slice_intervals( - total_rows=1000, + result = calc_slice_intervals( + [("/data/test.parquet", 1000)], worker_id=0, num_workers=1, + batch_size=1, checkpoint_state=checkpoint_state, - input_path="/data/test.parquet", - ) + )["/data/test.parquet"] self.assertEqual(result, [(100, 500), (600, 1000)]) def test_calc_slice_intervals_two_workers(self): @@ -148,21 +257,21 @@ def test_calc_slice_intervals_two_workers(self): } # Worker 0 gets first half of total rows - result_w0, _ = calc_slice_intervals( - total_rows=1000, + result_w0 = calc_slice_intervals( + [("/data/test.parquet", 1000)], worker_id=0, num_workers=2, + batch_size=1, checkpoint_state=checkpoint_state, - input_path="/data/test.parquet", - ) + )["/data/test.parquet"] # Worker 1 gets second half - result_w1, _ = calc_slice_intervals( - total_rows=1000, + result_w1 = calc_slice_intervals( + [("/data/test.parquet", 1000)], worker_id=1, num_workers=2, + batch_size=1, checkpoint_state=checkpoint_state, - input_path="/data/test.parquet", - ) + )["/data/test.parquet"] # Combined should cover all intervals total_rows_w0 = sum(end - start for start, end in result_w0) @@ -173,13 +282,13 @@ def test_calc_slice_intervals_empty_intervals(self): """Test calc_slice_intervals with empty intervals (fully consumed).""" # All data consumed: checkpoint at row 999 (last row) checkpoint_state = {"/data/test.parquet:0": 999} - result, _ = calc_slice_intervals( - total_rows=1000, + result = calc_slice_intervals( + [("/data/test.parquet", 1000)], worker_id=0, num_workers=2, + batch_size=1, checkpoint_state=checkpoint_state, - input_path="/data/test.parquet", - ) + )["/data/test.parquet"] self.assertEqual(result, []) def test_calc_slice_intervals_topology_change(self): @@ -194,13 +303,13 @@ def test_calc_slice_intervals_topology_change(self): # Now redistribute among 3 workers total_rows = 0 for worker_id in range(3): - result, _ = calc_slice_intervals( - total_rows=1000, + result = calc_slice_intervals( + [("/data/test.parquet", 1000)], worker_id=worker_id, num_workers=3, + batch_size=1, checkpoint_state=checkpoint_state, - input_path="/data/test.parquet", - ) + )["/data/test.parquet"] for start, end in result: total_rows += end - start diff --git a/tzrec/protos/data.proto b/tzrec/protos/data.proto index fb1af5d2a..81cccc2b3 100644 --- a/tzrec/protos/data.proto +++ b/tzrec/protos/data.proto @@ -156,6 +156,10 @@ message DataConfig { // whether dataloader batches are returned in first-in, first-out order optional bool in_order = 28 [default = true]; + // drop train/eval tail batches with fewer rows than this, 0 disables. + // drop_remainder=true is equivalent to min_batch_size=batch_size. + optional uint32 min_batch_size = 29 [default = 0]; + // negative sampler oneof sampler { NegativeSampler negative_sampler = 101; From 6a17bd9a49c6a1cfd443caa2a04903141033bd94 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Sun, 20 Sep 2026 09:43:50 +0800 Subject: [PATCH 2/8] [bugfix] pass world_size to to_batches instead of inferring it calc_slice_intervals read the world size from the process group and required the slice count to be a multiple of it, so a tool calling to_batches() under torchrun to read the whole table on every rank, as build_faiss_index does in the hitrate job, failed the assertion. The slice geometry is now an explicit optional world_size argument of to_batches: without it every slice is an independent even share as before, and only BaseDataset, which knows the dataloader workers of each rank, passes it. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01163rPXLGHiM74LxdbWjMox --- tzrec/datasets/csv_dataset.py | 2 +- tzrec/datasets/dataset.py | 22 +++++++++----- tzrec/datasets/dataset_test.py | 2 +- tzrec/datasets/kafka_dataset.py | 4 ++- tzrec/datasets/odps_dataset.py | 3 +- tzrec/datasets/odps_dataset_v1.py | 2 +- tzrec/datasets/parquet_dataset.py | 3 +- tzrec/datasets/parquet_dataset_test.py | 6 +++- tzrec/datasets/utils.py | 41 +++++++++++++------------- 9 files changed, 50 insertions(+), 35 deletions(-) diff --git a/tzrec/datasets/csv_dataset.py b/tzrec/datasets/csv_dataset.py index 917f4c165..2129164f2 100644 --- a/tzrec/datasets/csv_dataset.py +++ b/tzrec/datasets/csv_dataset.py @@ -154,7 +154,7 @@ def schema(self) -> pa.Schema: return self._schema def to_batches( - self, worker_id: int = 0, num_workers: int = 1 + self, worker_id: int = 0, num_workers: int = 1, world_size: Optional[int] = None ) -> Iterator[Dict[str, pa.Array]]: """Get batch iterator.""" input_files = self._input_files[worker_id::num_workers] diff --git a/tzrec/datasets/dataset.py b/tzrec/datasets/dataset.py index c7229c4bc..351031b54 100644 --- a/tzrec/datasets/dataset.py +++ b/tzrec/datasets/dataset.py @@ -300,8 +300,8 @@ def input_fields(self) -> List[pa.Field]: assert self._input_fields is not None return self._input_fields - def get_worker_info(self) -> Tuple[int, int]: - """Get multiprocessing dataloader worker id and worker number.""" + def get_worker_info(self) -> Tuple[int, int, int]: + """Get global dataloader worker id, worker number and world size.""" worker_info = get_worker_info() if worker_info is None: worker_id = 0 @@ -317,7 +317,7 @@ def get_worker_info(self) -> Tuple[int, int]: rank = 0 world_size = 1 - return rank * num_workers + worker_id, num_workers * world_size + return rank * num_workers + worker_id, num_workers * world_size, world_size def load_state_dict(self, state: Optional[Dict[str, Any]]) -> None: """Set checkpoint state for resume. @@ -332,8 +332,8 @@ def __iter__(self) -> Iterator[Batch]: if self._sampler is not None and not self._sampler_inited: self._sampler.init() self._sampler_inited = True - worker_id, num_workers = self.get_worker_info() - for input_data in self._reader.to_batches(worker_id, num_workers): + worker_id, num_workers, world_size = self.get_worker_info() + for input_data in self._reader.to_batches(worker_id, num_workers, world_size): yield self._build_batch(input_data) # pass complete: clear the resume state so later epochs do full passes self._reader.load_state_dict(None) @@ -601,9 +601,17 @@ def schema(self) -> pa.Schema: raise NotImplementedError def to_batches( - self, worker_id: int = 0, num_workers: int = 1 + self, worker_id: int = 0, num_workers: int = 1, world_size: Optional[int] = None ) -> Iterator[Dict[str, pa.Array]]: - """Get batch iterator.""" + """Get batch iterator of one of ``num_workers`` slices of the data. + + Args: + worker_id (int): slice id. + num_workers (int): slice number; without ``world_size`` every slice + is an independent even share. + world_size (int, optional): set by the dataset when the slices are + the rank-major dataloader workers of ``world_size`` ranks. + """ raise NotImplementedError def _slice_buff_data( diff --git a/tzrec/datasets/dataset_test.py b/tzrec/datasets/dataset_test.py index 91d3e280c..c6ddb20e1 100644 --- a/tzrec/datasets/dataset_test.py +++ b/tzrec/datasets/dataset_test.py @@ -133,7 +133,7 @@ def _reader(self) -> Iterator[pa.RecordBatch]: ) def to_batches( - self, worker_id: int = 0, num_workers: int = 1 + self, worker_id: int = 0, num_workers: int = 1, world_size: Optional[int] = None ) -> Iterator[Dict[str, pa.Array]]: yield from self._arrow_reader_iter(self._reader()) diff --git a/tzrec/datasets/kafka_dataset.py b/tzrec/datasets/kafka_dataset.py index c35579766..1a614a056 100644 --- a/tzrec/datasets/kafka_dataset.py +++ b/tzrec/datasets/kafka_dataset.py @@ -376,6 +376,7 @@ def _reader( Args: worker_id: Worker ID num_workers: Total number of workers + world_size: Unused, partitions are split by global worker id Yields: PyArrow RecordBatch @@ -627,13 +628,14 @@ def _heartbeat() -> None: logger.warning(f"consumer.close() failed: {e}") def to_batches( - self, worker_id: int = 0, num_workers: int = 1 + self, worker_id: int = 0, num_workers: int = 1, world_size: Optional[int] = None ) -> Iterator[Dict[str, pa.Array]]: """Get batch iterator. Args: worker_id: Worker ID num_workers: Total number of workers + world_size: Unused, partitions are split by global worker id Yields: Dict of column name to PyArrow Array diff --git a/tzrec/datasets/odps_dataset.py b/tzrec/datasets/odps_dataset.py index d992f1223..f08bdbdfa 100644 --- a/tzrec/datasets/odps_dataset.py +++ b/tzrec/datasets/odps_dataset.py @@ -617,7 +617,7 @@ def load_state_dict(self, state: Optional[Dict[str, int]]) -> None: self._restore_sessions(state) def to_batches( - self, worker_id: int = 0, num_workers: int = 1 + self, worker_id: int = 0, num_workers: int = 1, world_size: Optional[int] = None ) -> Iterator[Dict[str, pa.Array]]: """Get batch iterator.""" # (source_id_prefix, record_count, client, session) in read order @@ -642,6 +642,7 @@ def to_batches( self._equalize_rank_steps, self._min_batch_size, checkpoint_state=self._checkpoint_state, + world_size=world_size, ) def _combined_reader() -> Iterator[pa.RecordBatch]: diff --git a/tzrec/datasets/odps_dataset_v1.py b/tzrec/datasets/odps_dataset_v1.py index 6ffd5c947..efeec946d 100644 --- a/tzrec/datasets/odps_dataset_v1.py +++ b/tzrec/datasets/odps_dataset_v1.py @@ -176,7 +176,7 @@ def _iter_one_table( yield data def to_batches( - self, worker_id: int = 0, num_workers: int = 1 + self, worker_id: int = 0, num_workers: int = 1, world_size: Optional[int] = None ) -> Iterator[Dict[str, pa.Array]]: """Get batch iterator.""" for input_path in self._input_path.split(","): diff --git a/tzrec/datasets/parquet_dataset.py b/tzrec/datasets/parquet_dataset.py index 0d48712e9..3a5c104ad 100644 --- a/tzrec/datasets/parquet_dataset.py +++ b/tzrec/datasets/parquet_dataset.py @@ -240,7 +240,7 @@ def schema(self) -> pa.Schema: return self._schema def to_batches( - self, worker_id: int = 0, num_workers: int = 1 + self, worker_id: int = 0, num_workers: int = 1, world_size: Optional[int] = None ) -> Iterator[Dict[str, pa.Array]]: """Get batch iterator.""" if len(self._input_files) == 0: @@ -256,6 +256,7 @@ def to_batches( self._equalize_rank_steps, self._min_batch_size, checkpoint_state=self._checkpoint_state, + world_size=world_size, )[self._input_path] def _combined_reader() -> Iterator[pa.RecordBatch]: diff --git a/tzrec/datasets/parquet_dataset_test.py b/tzrec/datasets/parquet_dataset_test.py index 370bdc870..524a00fd2 100644 --- a/tzrec/datasets/parquet_dataset_test.py +++ b/tzrec/datasets/parquet_dataset_test.py @@ -537,7 +537,11 @@ def _reader_worker(rank, port, queue): min_batch_size=min_batch_size, equalize_rank_steps=True, ) - queue.put((rank, [len(b["id_a"]) for b in reader.to_batches(rank, 2)])) + sizes = [len(b["id_a"]) for b in reader.to_batches(rank, 2, world_size=2)] + # a tool slicing on its own terms, e.g. faiss_util.build_faiss_index, + # still reads the whole table on every rank + assert sum(len(b["id_a"]) for b in reader.to_batches()) == num_rows + queue.put((rank, sizes)) t = pa.Table.from_arrays( [pa.array(["1"] * num_rows), pa.array([0] * num_rows)], diff --git a/tzrec/datasets/utils.py b/tzrec/datasets/utils.py index b89a67d21..70d4cd1b5 100644 --- a/tzrec/datasets/utils.py +++ b/tzrec/datasets/utils.py @@ -10,7 +10,6 @@ # limitations under the License. import glob -import os import re from dataclasses import dataclass, field from typing import Dict, List, Optional, Tuple @@ -818,12 +817,6 @@ def calc_remaining_intervals( return remaining if remaining else [] -def _get_world_size() -> int: - if dist.is_initialized(): - return dist.get_world_size() - return int(os.environ.get("WORLD_SIZE", 1)) - - def plan_rank_worker_intervals( source_rows: List[int], rank: int, @@ -968,32 +961,38 @@ def calc_slice_intervals( equalize_rank_steps: bool = False, min_batch_size: int = 0, checkpoint_state: Optional[Dict[str, int]] = None, + world_size: Optional[int] = None, ) -> Dict[str, List[Tuple[int, int]]]: - """Assign the row intervals of every source to one global dataloader worker. + """Assign the row intervals of every source to one of ``num_workers`` slices. - ``worker_id`` and ``num_workers`` are the rank-major global ids from - ``BaseDataset.get_worker_info``. The rows each source still has after the - checkpoint state is applied are planned per rank and worker by - ``plan_rank_worker_intervals`` and mapped back onto the remaining intervals. + Without ``world_size`` every slice is an independent even share of the rows. + With ``world_size`` the slices are the rank-major dataloader workers of + ``BaseDataset.get_worker_info``, ``num_workers // world_size`` per rank, and + the rows each source still has after the checkpoint state is applied are + planned per rank and worker by ``plan_rank_worker_intervals`` and mapped back + onto the remaining intervals. Args: sources (list): (source_id_prefix, total_rows) in read order. - worker_id (int): global worker id. - num_workers (int): total worker number over all ranks. + worker_id (int): slice id. + num_workers (int): slice number. batch_size (int): batch size. equalize_rank_steps (bool): make every rank yield the same batch sizes. min_batch_size (int): drop a final batch with fewer rows, 0 disables. checkpoint_state (dict): dict mapping source_id to max consumed row index. + world_size (int, optional): number of ranks the slices are grouped into. Returns: - dict mapping source_id_prefix to the (start, end) intervals of this worker. + dict mapping source_id_prefix to the (start, end) intervals of this slice. """ - world_size = _get_world_size() - assert num_workers % world_size == 0, ( - f"num_workers[{num_workers}] must be a multiple of world_size[{world_size}]" - ) - local_workers = num_workers // world_size - rank, local_worker_id = divmod(worker_id, local_workers) + if world_size is None: + rank, world_size, local_worker_id, local_workers = worker_id, num_workers, 0, 1 + else: + assert num_workers % world_size == 0, ( + f"num_workers[{num_workers}] must be a multiple of world_size[{world_size}]" + ) + local_workers = num_workers // world_size + rank, local_worker_id = divmod(worker_id, local_workers) remaining = [ calc_remaining_intervals(checkpoint_state, prefix, total_rows) From 05d69c371abe827ee01fa0b29dce99b595038b34 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Sun, 20 Sep 2026 09:50:15 +0800 Subject: [PATCH 3/8] [bugfix] apply drop_remainder and min_batch_size to training only Eval metrics and predict output must cover every row, and BatchNorm runs on running statistics outside training, so a short final batch is harmless there. Both tail-drop knobs now take effect only in Mode.TRAIN; rank step equalization still applies to train and eval. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01163rPXLGHiM74LxdbWjMox --- docs/source/feature/data.md | 4 ++-- tzrec/datasets/dataset.py | 10 +++++----- tzrec/datasets/parquet_dataset_test.py | 14 ++++++++++---- tzrec/protos/data.proto | 4 ++-- 4 files changed, 19 insertions(+), 13 deletions(-) diff --git a/docs/source/feature/data.md b/docs/source/feature/data.md index 51ff801ca..cc67ff808 100644 --- a/docs/source/feature/data.md +++ b/docs/source/feature/data.md @@ -330,11 +330,11 @@ pipeline.global-job-parameters: | ### drop_remainder - 是否丢弃掉最后一个不足batch_size的batch数据,默认为false -- 仅在训练和评估时生效,预测时不会丢弃任何数据;等价于`min_batch_size`设置为`batch_size` +- 仅在训练时生效,评估和预测时不会丢弃任何数据;等价于`min_batch_size`设置为`batch_size` ### min_batch_size -- 训练和评估时,丢弃掉每个数据读取进程最后一个行数小于`min_batch_size`的batch,默认为0(不丢弃) +- 训练时,丢弃掉每个数据读取进程最后一个行数小于`min_batch_size`的batch,默认为0(不丢弃);评估和预测时不生效 - 使用BatchNorm等要求batch内至少2行样本的模型时,建议设置为2,避免样本表行数恰好使最后一个batch只剩1行导致训练失败 - 注:OdpsDataset和ParquetDataset按行切分数据时,会保证每个`proc`(rank)读取相同的步数,以避免同步训练时卡住;为此每张表(或每个分区)最多有`nproc - 1`行样本不会被读取,预测时不受影响 diff --git a/tzrec/datasets/dataset.py b/tzrec/datasets/dataset.py index 351031b54..ce330d8d9 100644 --- a/tzrec/datasets/dataset.py +++ b/tzrec/datasets/dataset.py @@ -208,11 +208,8 @@ def __init__( self._batch_size = config_util.get_inference_batch_size(data_config) else: self._batch_size = data_config.batch_size - # predict keeps every row: no tail dropping and no rank equalization - if mode == Mode.PREDICT: - self._drop_remainder = False - self._min_batch_size = 0 - else: + # only training drops tail batches; eval and predict keep every row + if mode == Mode.TRAIN: self._drop_remainder = data_config.drop_remainder self._min_batch_size = data_config.min_batch_size if self._min_batch_size > self._batch_size: @@ -220,6 +217,9 @@ def __init__( f"data_config.min_batch_size[{self._min_batch_size}] must not " f"exceed the batch size[{self._batch_size}]." ) + else: + self._drop_remainder = False + self._min_batch_size = 0 self._sampler = None self._sampler_inited = False diff --git a/tzrec/datasets/parquet_dataset_test.py b/tzrec/datasets/parquet_dataset_test.py index 524a00fd2..f6241ef1b 100644 --- a/tzrec/datasets/parquet_dataset_test.py +++ b/tzrec/datasets/parquet_dataset_test.py @@ -344,7 +344,10 @@ def test_create_dataloader_tail_batches(self, min_batch_size, expected): ) self.assertEqual(sizes, expected) - def test_create_dataloader_predict_keeps_every_row(self): + @parameterized.expand( + [[Mode.EVAL], [Mode.PREDICT]], name_func=parameterized_name_func + ) + def test_create_dataloader_eval_predict_keep_every_row(self, mode): feature_cfgs = self._create_feature_cfgs() features = create_features(feature_cfgs) with tempfile.TemporaryDirectory(prefix="tzrec_") as test_dir: @@ -363,10 +366,13 @@ def test_create_dataloader_predict_keeps_every_row(self): features, f"{test_dir}/*", reserved_columns=["label"], - mode=Mode.PREDICT, + mode=mode, ) num_rows = sum( - len(batch.reserves.get()) for batch in dataloader.get_iterator() + len(batch.reserves.get()) + if mode == Mode.PREDICT + else len(batch.labels["label"]) + for batch in dataloader.get_iterator() ) self.assertEqual(num_rows, 8201) @@ -383,7 +389,7 @@ def test_min_batch_size_exceeds_batch_size(self): min_batch_size=5, ) with self.assertRaisesRegex(ValueError, "min_batch_size"): - ParquetDataset(data_config, features, f"{test_dir}/*") + ParquetDataset(data_config, features, f"{test_dir}/*", mode=Mode.TRAIN) def test_create_dataloader_get_iterator_reuses_eager_iterator(self): """get_iterator() yields the eagerly prefetched iterator first. diff --git a/tzrec/protos/data.proto b/tzrec/protos/data.proto index 81cccc2b3..aa391336f 100644 --- a/tzrec/protos/data.proto +++ b/tzrec/protos/data.proto @@ -104,7 +104,7 @@ message DataConfig { // mini batch size to use for and evaluation. optional uint32 eval_batch_size = 11; - // drop last batch less than batch_size + // drop last training batch less than batch_size optional bool drop_remainder = 12 [default = false]; // fg threads for each worker, @@ -156,7 +156,7 @@ message DataConfig { // whether dataloader batches are returned in first-in, first-out order optional bool in_order = 28 [default = true]; - // drop train/eval tail batches with fewer rows than this, 0 disables. + // drop training tail batches with fewer rows than this, 0 disables. // drop_remainder=true is equivalent to min_batch_size=batch_size. optional uint32 min_batch_size = 29 [default = 0]; From 91e147d52035feae90e5d98536950f8a56f5994a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Sun, 20 Sep 2026 11:52:55 +0800 Subject: [PATCH 4/8] [bugfix] plan a rank's sources as one stream and cache odps session row counts The planner laid every table or session out on its own: each source was split by world_size, losing up to world_size - 1 rows per source, and each source could leave a partial batch even though a worker's buffer carries its tail into the next source, so a multi-source pass ended with one partial batch per source rather than per rank. The sources of a rank are now planned as one stream: ranks split it as if rows were dealt round-robin, so every source stays spread over all ranks and the rank totals differ by at most one row over the whole pass; inside a rank the stream is cut into whole batches plus one tail, dealt source by source to the least loaded workers, with the rows completing the open batch going to its holder. Extra rows and the min_batch_size cut are decided once per pass on the stream length. ODPS session record counts are fetched once when the sessions are created or restored and cached on the reader instead of by every dataloader worker at the start of every epoch. calc_slice_intervals returns a list aligned with its sources so duplicate input paths no longer collapse to one plan, and the two-rank parquet test asserts inside the children instead of blocking on a queue. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01163rPXLGHiM74LxdbWjMox --- docs/source/feature/data.md | 5 +- tzrec/datasets/csv_dataset.py | 3 +- tzrec/datasets/dataset.py | 4 +- tzrec/datasets/odps_dataset.py | 28 +++- tzrec/datasets/parquet_dataset.py | 4 +- tzrec/datasets/parquet_dataset_test.py | 36 +++-- tzrec/datasets/utils.py | 196 ++++++++++++++----------- tzrec/datasets/utils_test.py | 115 +++++++++++---- 8 files changed, 249 insertions(+), 142 deletions(-) diff --git a/docs/source/feature/data.md b/docs/source/feature/data.md index cc67ff808..3b901be21 100644 --- a/docs/source/feature/data.md +++ b/docs/source/feature/data.md @@ -33,6 +33,8 @@ data_config { - `odps://{project}/tables/{table_name}/{partition}`,多表按逗号分隔 - 如果单表需要设置多个分区,可以用`&`简写,来分隔多个分区,`odps://{project}/tables/{table_name}/{partition1}&{partition2}` +- 注意: 训练和评估时,OdpsDataset会保证每个`proc`读取相同的步数以避免同步训练卡住,为此全部输入合计最多有`nproc - 1`行样本不会被读取;预测时不受影响 + - 运行训练/评估/导出/预测等命令时 - **本地环境**: @@ -75,6 +77,8 @@ data_config { - 注意: 如果每个parquet文件中的数据量不相等或文件数据小于worker数,ParquetDataset会自动重分配数据,来保证每个worker读取的数据量相等。但仍建议parquet文件数是 `nproc-per-node * nnodes * num_workers`的倍数,并且每个parquet文件的数据量基本相等,减少数据自动重分配的IO开销。 +- 注意: 训练和评估时,ParquetDataset会保证每个`proc`读取相同的步数以避免同步训练卡住,为此全部输入合计最多有`nproc - 1`行样本不会被读取;预测时不受影响 + ## CsvDataset - input_path: 按如下格式设置 @@ -336,7 +340,6 @@ pipeline.global-job-parameters: | - 训练时,丢弃掉每个数据读取进程最后一个行数小于`min_batch_size`的batch,默认为0(不丢弃);评估和预测时不生效 - 使用BatchNorm等要求batch内至少2行样本的模型时,建议设置为2,避免样本表行数恰好使最后一个batch只剩1行导致训练失败 -- 注:OdpsDataset和ParquetDataset按行切分数据时,会保证每个`proc`(rank)读取相同的步数,以避免同步训练时卡住;为此每张表(或每个分区)最多有`nproc - 1`行样本不会被读取,预测时不受影响 ### batch_cost_size diff --git a/tzrec/datasets/csv_dataset.py b/tzrec/datasets/csv_dataset.py index 2129164f2..8ac99ff94 100644 --- a/tzrec/datasets/csv_dataset.py +++ b/tzrec/datasets/csv_dataset.py @@ -87,13 +87,14 @@ class CsvReader(BaseReader): input_path (str): data input path. batch_size (int): batch size. selected_cols (list): selection column names. - drop_remainder (bool): drop last batch. + drop_remainder (bool): drop last batch, same as min_batch_size=batch_size. shuffle (bool): shuffle data or not. shuffle_buffer_size (int): buffer size for shuffle. column_names (list): set column name if csv without header. delimiter (str): csv delimiter. sample_cost_field (str): sample cost field name. batch_cost_size (int): batch cost limit size. + min_batch_size (int): drop a final batch with fewer rows, 0 disables. """ def __init__( diff --git a/tzrec/datasets/dataset.py b/tzrec/datasets/dataset.py index ce330d8d9..028b4d23d 100644 --- a/tzrec/datasets/dataset.py +++ b/tzrec/datasets/dataset.py @@ -551,8 +551,8 @@ class BaseReader(metaclass=_reader_meta_cls): sample_cost_field (str): sample cost field name. batch_cost_size (int): batch cost limit size. min_batch_size (int): drop a final batch with fewer rows, 0 disables. - equalize_rank_steps (bool): make every rank yield the same batch sizes, - honored by readers that slice rows by count. + equalize_rank_steps (bool): make every rank yield the same number of + batches, honored by readers that slice rows by count. """ def __init__( diff --git a/tzrec/datasets/odps_dataset.py b/tzrec/datasets/odps_dataset.py index f08bdbdfa..22ccaf9dd 100644 --- a/tzrec/datasets/odps_dataset.py +++ b/tzrec/datasets/odps_dataset.py @@ -15,6 +15,7 @@ import threading import time from collections import OrderedDict +from concurrent.futures import ThreadPoolExecutor from typing import Any, Dict, Iterator, List, Optional, Tuple, Union import pyarrow as pa @@ -414,7 +415,7 @@ class OdpsReader(BaseReader): sample_cost_field (str): sample cost field name. batch_cost_size (int): batch cost limit size. min_batch_size (int): drop a final batch with fewer rows, 0 disables. - equalize_rank_steps (bool): make every rank yield the same batch sizes. + equalize_rank_steps (bool): make every rank yield the same number of batches. """ def __init__( @@ -453,6 +454,7 @@ def __init__( self._proj_to_o = {} self._table_to_cli = {} self._input_to_sess = {} + self._sess_row_counts: Dict[str, int] = {} self._init_client() fields = [] @@ -520,11 +522,22 @@ def _init_session(self) -> None: else: session_ids.append(None) + # a session's record count is an immutable snapshot: fetch it once + # here instead of in every dataloader worker + sess_infos = [(x, None) for x in session_ids] + if int(os.environ.get("RANK", 0)) == 0: + sess_reqs = [SessionRequest(session_id=x) for x in session_ids] + with ThreadPoolExecutor(max_workers=8) as executor: + record_counts = executor.map( + _get_session_record_count, [client] * len(sess_reqs), sess_reqs + ) + sess_infos = list(zip(session_ids, record_counts)) if self._pg is not None: - dist.broadcast_object_list(session_ids, group=self._pg) + dist.broadcast_object_list(sess_infos, group=self._pg) self._input_to_sess[input_path] = [ - SessionRequest(session_id=x) for x in session_ids + SessionRequest(session_id=x) for x, _ in sess_infos ] + self._sess_row_counts.update(dict(sess_infos)) # refresh session if int(os.environ.get("RANK", 0)) == 0: t = threading.Thread( @@ -586,6 +599,7 @@ def _restore_sessions(self, checkpoint_state: Dict[str, int]) -> None: " Please restart training from scratch." ) restored_sess_reqs.append(sess_req) + self._sess_row_counts[session_id] = resp.record_count except ODPSError as e: raise RuntimeError( f"Cannot resume from checkpoint: ODPS session {session_id} " @@ -629,12 +643,12 @@ def to_batches( sources.append( ( f"{input_path}#{sess_req.session_id}", - _get_session_record_count(client, sess_req), + self._sess_row_counts[sess_req.session_id], client, sess_req, ) ) - intervals = calc_slice_intervals( + plan = calc_slice_intervals( [(prefix, record_count) for prefix, record_count, _, _ in sources], worker_id, num_workers, @@ -646,8 +660,8 @@ def to_batches( ) def _combined_reader() -> Iterator[pa.RecordBatch]: - for prefix, _, client, sess_req in sources: - for start, end in intervals[prefix]: + for (prefix, _, client, sess_req), intervals in zip(sources, plan): + for start, end in intervals: if start >= end: continue yield from _reader_iter( diff --git a/tzrec/datasets/parquet_dataset.py b/tzrec/datasets/parquet_dataset.py index 3a5c104ad..8a1cef6ee 100644 --- a/tzrec/datasets/parquet_dataset.py +++ b/tzrec/datasets/parquet_dataset.py @@ -161,7 +161,7 @@ class ParquetReader(BaseReader): sample_cost_field (str): sample cost field name. batch_cost_size (int): batch cost limit size. min_batch_size (int): drop a final batch with fewer rows, 0 disables. - equalize_rank_steps (bool): make every rank yield the same batch sizes. + equalize_rank_steps (bool): make every rank yield the same number of batches. """ def __init__( @@ -257,7 +257,7 @@ def to_batches( self._min_batch_size, checkpoint_state=self._checkpoint_state, world_size=world_size, - )[self._input_path] + )[0] def _combined_reader() -> Iterator[pa.RecordBatch]: for start, end in worker_intervals: diff --git a/tzrec/datasets/parquet_dataset_test.py b/tzrec/datasets/parquet_dataset_test.py index f6241ef1b..4daabc5af 100644 --- a/tzrec/datasets/parquet_dataset_test.py +++ b/tzrec/datasets/parquet_dataset_test.py @@ -517,21 +517,23 @@ def _reader_worker(rank, port): @parameterized.expand( [ - # extra row would buy rank 0 a fifth step: dropped - [33, 4, 0, [4, 4], [[4] * 4, [4] * 4]], - # residue-1 tails stay when min_batch_size is 0 - [35, 4, 0, [5, 5], [[4] * 4 + [2], [4] * 4 + [1]]], + # one worker per rank; extra row would buy one rank a fifth step: dropped + [33, 4, 0, 1, [[4] * 4, [4] * 4]], + # residue-1 tails stay when min_batch_size is 0 (rank 0 holds the extra) + [35, 4, 0, 1, [[4] * 4 + [2], [4] * 4 + [1]]], # ... and go together with the extra row when min_batch_size=2 - [35, 4, 2, [4, 4], [[4] * 4, [4] * 4]], + [35, 4, 2, 1, [[4] * 4, [4] * 4]], # 8200 rows over two ranks: 4100 each, 4-row tails - [8200, 1024, 2, [5, 5], [[1024] * 4 + [4], [1024] * 4 + [4]]], + [8200, 1024, 2, 1, [[1024] * 4 + [4], [1024] * 4 + [4]]], + # two workers per rank: full batches dealt two per worker, tail last + [35, 4, 0, 2, [[4] * 4 + [2], [4] * 4 + [1]]], ], name_func=parameterized_name_func, ) def test_parquet_reader_equal_steps( - self, num_rows, batch_size, min_batch_size, steps, sizes + self, num_rows, batch_size, min_batch_size, local_workers, sizes ): - def _reader_worker(rank, port, queue): + def _reader_worker(rank, port): os.environ["RANK"] = str(rank) os.environ["WORLD_SIZE"] = str(2) os.environ["MASTER_ADDR"] = "127.0.0.1" @@ -543,11 +545,18 @@ def _reader_worker(rank, port, queue): min_batch_size=min_batch_size, equalize_rank_steps=True, ) - sizes = [len(b["id_a"]) for b in reader.to_batches(rank, 2, world_size=2)] + rank_sizes = [] + for w in range(local_workers): + rank_sizes += [ + len(b["id_a"]) + for b in reader.to_batches( + rank * local_workers + w, 2 * local_workers, world_size=2 + ) + ] + assert sorted(rank_sizes, reverse=True) == sizes[rank], rank_sizes # a tool slicing on its own terms, e.g. faiss_util.build_faiss_index, # still reads the whole table on every rank assert sum(len(b["id_a"]) for b in reader.to_batches()) == num_rows - queue.put((rank, sizes)) t = pa.Table.from_arrays( [pa.array(["1"] * num_rows), pa.array([0] * num_rows)], @@ -560,20 +569,15 @@ def _reader_worker(rank, port, queue): writer.close() port = misc_util.get_free_port() - queue = mp.Queue() procs = [ - mp.Process(target=_reader_worker, args=(rank, port, queue)) - for rank in range(2) + mp.Process(target=_reader_worker, args=(rank, port)) for rank in range(2) ] for p in procs: p.start() - results = dict(queue.get() for _ in procs) for i, p in enumerate(procs): p.join() if p.exitcode != 0: raise RuntimeError(f"reader worker-{i} failed.") - self.assertEqual([len(results[r]) for r in range(2)], steps) - self.assertEqual([results[r] for r in range(2)], sizes) class ParquetWriterTest(unittest.TestCase): diff --git a/tzrec/datasets/utils.py b/tzrec/datasets/utils.py index 70d4cd1b5..a451a3a9f 100644 --- a/tzrec/datasets/utils.py +++ b/tzrec/datasets/utils.py @@ -826,24 +826,28 @@ def plan_rank_worker_intervals( batch_size: int, equalize_rank_steps: bool = False, min_batch_size: int = 0, -) -> List[Tuple[int, int]]: - """Plan the row range one dataloader worker reads from every source. +) -> List[List[Tuple[int, int]]]: + """Plan the row intervals one dataloader worker reads from every source. Sources (tables, sessions, or remaining checkpoint intervals) are consumed in - order by a single buffered reader per worker, so a worker's partial batches - are decided by its cumulative row count over all sources. The plan is - computed identically on every rank: - - 1. Every rank takes ``rows // world_size`` rows of each source, laid out - contiguously; the ``rows % world_size`` extra rows are read one each by - the lowest ranks. - 2. Inside a rank, whole batches are spread over its workers (rotating the - first worker across sources) and the partial tail goes to the worker laid - out last, so a worker never manufactures a partial batch of its own. - 3. A cumulative tail smaller than ``min_batch_size`` is dropped. - 4. With ``equalize_rank_steps`` every rank yields the same batch-size list: - the extra rows landing on a worker are kept only when they cannot push - its last batch past ``batch_size``. + order by a single buffered reader per worker, so they are planned as one + stream, identically on every rank: + + 1. Ranks split the stream as if its rows were dealt round-robin: of the + first ``S`` rows rank ``r`` owns ``(S + W - 1 - r) // W``, laid out + contiguously inside every source, so every source stays spread over all + ranks (partition order is kept) while the rank totals are ``R // W`` or + ``R // W + 1``. + 2. With ``equalize_rank_steps`` every rank yields the same number of batches: + a rank drops its extra row only when it would buy a step, that is when + ``(R // W) % batch_size == 0``. A final batch shorter than + ``min_batch_size`` is cut off the end of the stream, on every rank alike. + 3. Inside a rank the stream is cut into whole batches plus one tail and dealt + to the workers source by source: the rows completing the open batch go to + its holder, whole batches are dealt as contiguous blocks to the least + loaded workers first, and the tail opens the next batch on the least + loaded worker. Only the holder ever buffers a partial batch, so a rank + ends a pass with at most one. Dropped rows are never read. @@ -854,79 +858,86 @@ def plan_rank_worker_intervals( worker_id (int): dataloader worker id within the rank. num_workers (int): dataloader workers per rank. batch_size (int): batch size. - equalize_rank_steps (bool): make every rank yield the same batch sizes. + equalize_rank_steps (bool): make every rank yield the same number of batches. min_batch_size (int): drop a final batch with fewer rows, 0 disables. Returns: - (start, end) per source for this rank and worker, ``start == end`` when - the worker reads nothing from that source. + for every source, the (start, end) intervals this worker reads, in read + order; empty when the worker reads nothing from that source. """ assert 0 < num_workers and 0 <= worker_id < num_workers assert 0 < world_size and 0 <= rank < world_size assert 0 <= min_batch_size <= batch_size - # chunk layout of each source inside a rank: (worker, rows) in layout order - layouts: List[List[Tuple[int, int]]] = [] - last_workers: List[int] = [] - cursor = 0 + def _rows_before(num_rows: int, r: int) -> int: + # rows of ranks < r among the first num_rows rows of the stream + return num_rows // world_size * r + min(num_rows % world_size, r) + + # this rank's (start, rows) in every source + shares: List[Tuple[int, int]] = [] + prefix = 0 for rows in source_rows: - full, tail = divmod(rows // world_size, batch_size) - order = [(cursor + i) % num_workers for i in range(num_workers)] - q, r = divmod(full, num_workers) - layout = [ - (w, (q + 1 if i < r else q) * batch_size) for i, w in enumerate(order) - ] - layout[-1] = (order[-1], layout[-1][1] + tail) - layouts.append(layout) - last_workers.append(order[-1]) - cursor = (cursor + full + (1 if tail else 0)) % num_workers - - totals = [0] * num_workers - for layout in layouts: - for w, cnt in layout: - totals[w] += cnt - - # drop a cumulative tail shorter than min_batch_size from the end of the - # worker's stream, walking its chunks backwards - shrink: Dict[Tuple[int, int], int] = {} - for w in range(num_workers): - remain = totals[w] % batch_size - if 0 < remain < min_batch_size: - totals[w] -= remain - for t in reversed(range(len(source_rows))): - cnt = dict(layouts[t])[w] - take = min(cnt, remain) - if take: - shrink[(t, w)] = take - remain -= take - if remain == 0: - break - - extras = [rows % world_size for rows in source_rows] + lo = _rows_before(prefix + rows, rank) - _rows_before(prefix, rank) + hi = _rows_before(prefix + rows, rank + 1) - _rows_before(prefix, rank + 1) + shares.append((lo, hi - lo)) + prefix += rows + total = sum(rows for _, rows in shares) + base = prefix // world_size + + cut = 0 if equalize_rank_steps: - landing = [0] * num_workers - for t, e in enumerate(extras): - if e: - landing[last_workers[t]] += 1 - for t, e in enumerate(extras): - w = last_workers[t] - remain = totals[w] % batch_size - if e and (remain == 0 or remain + landing[w] > batch_size): - extras[t] = 0 - - result = [] - for t, rows in enumerate(source_rows): - n = rows // world_size - pos = rank * n + min(rank, extras[t]) - start = end = pos - for i, (w, cnt) in enumerate(layouts[t]): - if i == num_workers - 1 and rank < extras[t]: - cnt += 1 - if w == worker_id: - start = pos - end = pos + cnt - shrink.get((t, w), 0) - pos += cnt - result.append((start, end)) + if base % batch_size == 0: + cut = total - base + elif base % batch_size < min_batch_size: + cut = total - base + base % batch_size + elif 0 < total % batch_size < min_batch_size: + cut = total % batch_size + for t in range(len(shares) - 1, -1, -1): + if cut == 0: + break + start, rows = shares[t] + take = min(rows, cut) + shares[t] = (start, rows - take) + cut -= take + + result: List[List[Tuple[int, int]]] = [] + loads = [0] * num_workers + holder, carry = 0, 0 + for start, rows in shares: + chunks: List[Tuple[int, int, int]] = [] + pos = 0 + if carry and rows: + pos = min(batch_size - carry, rows) + chunks.append((holder, 0, pos)) + loads[holder] += pos + carry = (carry + pos) % batch_size + full, tail = divmod(rows - pos, batch_size) + q, r = divmod(full, num_workers) + order = sorted(range(num_workers), key=lambda w: (loads[w], w)) + counts = {w: (q + 1 if i < r else q) * batch_size for i, w in enumerate(order)} + for w in order: + loads[w] += counts[w] + # the tail opens the next batch on the least loaded worker, whose block + # is laid out last so that block and tail form one interval + owner = min(range(num_workers), key=lambda w: (loads[w], w)) + for w in order: + if counts[w] and (w != owner or not tail): + chunks.append((w, pos, pos + counts[w])) + pos += counts[w] + if tail: + chunks.append((owner, pos, pos + counts[owner] + tail)) + loads[owner] += tail + holder, carry = owner, tail + + mine: List[Tuple[int, int]] = [] + for w, lo, hi in chunks: + if w != worker_id: + continue + if mine and mine[-1][1] == start + lo: + mine[-1] = (mine[-1][0], start + hi) + else: + mine.append((start + lo, start + hi)) + result.append(mine) return result @@ -962,7 +973,7 @@ def calc_slice_intervals( min_batch_size: int = 0, checkpoint_state: Optional[Dict[str, int]] = None, world_size: Optional[int] = None, -) -> Dict[str, List[Tuple[int, int]]]: +) -> List[List[Tuple[int, int]]]: """Assign the row intervals of every source to one of ``num_workers`` slices. Without ``world_size`` every slice is an independent even share of the rows. @@ -977,13 +988,14 @@ def calc_slice_intervals( worker_id (int): slice id. num_workers (int): slice number. batch_size (int): batch size. - equalize_rank_steps (bool): make every rank yield the same batch sizes. + equalize_rank_steps (bool): make every rank yield the same number of batches. min_batch_size (int): drop a final batch with fewer rows, 0 disables. checkpoint_state (dict): dict mapping source_id to max consumed row index. world_size (int, optional): number of ranks the slices are grouped into. Returns: - dict mapping source_id_prefix to the (start, end) intervals of this slice. + for every source, in the order of ``sources``, the (start, end) intervals + of this slice. """ if world_size is None: rank, world_size, local_worker_id, local_workers = worker_id, num_workers, 0, 1 @@ -1008,10 +1020,22 @@ def calc_slice_intervals( equalize_rank_steps, min_batch_size, ) - return { - prefix: _map_logical_range(intervals, start, end) - for (prefix, _), intervals, (start, end) in zip(sources, remaining, plan) - } + result = [ + [ + physical + for start, end in logical + for physical in _map_logical_range(intervals, start, end) + ] + for intervals, logical in zip(remaining, plan) + ] + num_rows = sum(end - start for intervals in result for start, end in intervals) + if equalize_rank_steps and min_batch_size < 2 and num_rows % batch_size == 1: + logger.warning( + "The final training batch of this pass has a single row; set " + "data_config.min_batch_size >= 2 if the model needs more than one " + "row per batch, e.g. with BatchNorm." + ) + return result def remove_nullable(field_type: pa.DataType) -> pa.DataType: diff --git a/tzrec/datasets/utils_test.py b/tzrec/datasets/utils_test.py index ac73c6dc9..d51a3a0ed 100644 --- a/tzrec/datasets/utils_test.py +++ b/tzrec/datasets/utils_test.py @@ -42,15 +42,19 @@ def _rank_batches( batch_size: int, equalize: bool, min_batch_size: int, - ) -> Tuple[List[List[int]], int]: - """Simulate every worker's buffered stream; return per-rank sizes, rows read.""" + ) -> Tuple[List[List[int]], int, List[int]]: + """Simulate every worker's buffered stream. + + Returns per-rank batch sizes, rows read, and per-worker row totals. + """ per_rank = [] + worker_totals = [] num_read = 0 seen = [set() for _ in rows] for rank in range(world_size): batches = [] for worker in range(num_workers): - ranges = plan_rank_worker_intervals( + intervals = plan_rank_worker_intervals( rows, rank, world_size, @@ -61,17 +65,20 @@ def _rank_batches( min_batch_size, ) total = 0 - for t, (start, end) in enumerate(ranges): - assert 0 <= start <= end <= rows[t] - assert seen[t].isdisjoint(range(start, end)) - seen[t].update(range(start, end)) - total += end - start + for t, source_intervals in enumerate(intervals): + for start, end in source_intervals: + assert 0 <= start < end <= rows[t] + assert seen[t].isdisjoint(range(start, end)) + seen[t].update(range(start, end)) + total += end - start num_read += total + worker_totals.append(total) + # a worker's stream is whole batches plus at most one final tail batches.extend([batch_size] * (total // batch_size)) if total % batch_size >= max(min_batch_size, 1): batches.append(total % batch_size) - per_rank.append(sorted(batches)) - return per_rank, num_read + per_rank.append(batches) + return per_rank, num_read, worker_totals def test_plan_rank_worker_intervals_invariants(self): rng = random.Random(0) @@ -81,12 +88,12 @@ def test_plan_rank_worker_intervals_invariants(self): for min_batch_size in sorted({0, 1, min(2, batch_size), batch_size}): cases = [ [rng.randrange(3 * batch_size * world_size + 5) for _ in range(n)] - for n in (1, 2, 3) - for _ in range(20) + for n in (1, 2, 3, 5) + for _ in range(15) ] cases += [[r] for r in range(4 * batch_size * world_size + 3)] for rows in cases: - per_rank, num_read = self._rank_batches( + per_rank, num_read, _ = self._rank_batches( rows, world_size, num_workers, @@ -101,36 +108,55 @@ def test_plan_rank_worker_intervals_invariants(self): for batches in per_rank: self.assertTrue(all(b <= batch_size for b in batches), msg) self.assertTrue(all(b >= min_batch_size for b in batches), msg) + # at most one partial batch per rank and pass + self.assertLessEqual( + sum(1 for b in batches if b != batch_size), 1, msg + ) if equalize: self.assertEqual(len({len(b) for b in per_rank}), 1, msg) if not equalize and min_batch_size == 0: self.assertEqual(num_read, sum(rows), msg) else: - max_drop = (world_size - 1) * len(rows) - max_drop += ( - max(min_batch_size - 1, 0) * num_workers * world_size - ) + # the extra rows of the pass, plus a short tail per rank + max_drop = world_size - 1 + max_drop += max(min_batch_size - 1, 0) * world_size self.assertLessEqual(sum(rows) - num_read, max_drop, msg) + def test_plan_rank_worker_intervals_uniform_sources_balance(self): + # dealing to the least loaded worker keeps the skew independent of the + # number of sources + for num_workers, num_sources in ((2, 2), (2, 40), (4, 4), (4, 40), (8, 40)): + _, _, worker_totals = self._rank_batches( + [13] * num_sources, 1, num_workers, 4, True, 0 + ) + self.assertLessEqual(max(worker_totals) - min(worker_totals), 2 * 4) + @parameterized.expand( [ # the failing job: 8200 rows, 8 workers -> one 8-row tail, no drop [[8200], 1, 8, 1024, 0, 9, [8], 8200], [[8200], 1, 8, 1024, 2, 9, [8], 8200], [[8201], 1, 8, 1024, 2, 9, [9], 8201], - # extra row would buy a whole step on rank 0 -> drop it + # extra row would buy a whole step on one rank -> drop it [[33], 2, 4, 4, 0, 4, [], 32], # residue-1 tails are kept without min_batch_size [[34], 2, 4, 4, 0, 5, [1], 34], [[35], 2, 4, 4, 0, 5, [2], 35], - # ... and dropped together with the extras when min_batch_size=2 + # ... and dropped together with the extra row when min_batch_size=2 [[35], 2, 4, 4, 2, 4, [], 32], [[36], 2, 4, 4, 2, 5, [2], 36], # too few rows for one 2-row batch per rank -> empty pass [[3], 2, 4, 4, 2, 0, [], 0], - # tails of two sources (1000 + 25) sum to 1024 + 1 on one worker + # tails of two sources carry across the boundary: 2049 = 2 * 1024 + 1 [[1000, 1049], 1, 1, 1024, 0, 3, [1], 2049], [[1000, 1049], 1, 1, 1024, 2, 2, [], 2048], + # sources are one stream: 300 rows are exactly three batches + [[150, 150], 1, 4, 100, 0, 3, [], 300], + [[6, 7], 1, 3, 4, 0, 4, [1], 13], + [[13, 13], 1, 2, 4, 0, 7, [2], 26], + # ranks split the stream too: 80 rows over 8 ranks lose nothing + [[33, 47], 8, 1, 10, 0, 1, [], 80], + [[33, 47], 8, 2, 10, 2, 1, [], 80], # drop_remainder: no partial batch at all [[8200], 1, 8, 1024, 1024, 8, [], 8192], ], @@ -147,7 +173,7 @@ def test_plan_rank_worker_intervals( rank0_tails, num_read, ): - per_rank, actual_read = self._rank_batches( + per_rank, actual_read, _ = self._rank_batches( rows, world_size, num_workers, batch_size, True, min_batch_size ) self.assertEqual([len(b) for b in per_rank], [steps] * world_size) @@ -155,10 +181,45 @@ def test_plan_rank_worker_intervals( self.assertEqual(actual_read, num_read) def test_plan_rank_worker_intervals_predict_keeps_rows(self): - per_rank, num_read = self._rank_batches([35], 2, 4, 4, False, 0) + per_rank, num_read, _ = self._rank_batches([35], 2, 4, 4, False, 0) self.assertEqual([len(b) for b in per_rank], [5, 5]) self.assertEqual(num_read, 35) + def test_calc_slice_intervals_two_ranks_two_workers(self): + # 35 rows: rank 0 owns [0, 18) with the extra row, rank 1 owns [18, 35); + # inside each rank the 4 full batches are dealt two per worker and the + # tail opens on worker 0, whose block is laid out last + expected = {0: [(8, 18)], 1: [(0, 8)], 2: [(26, 35)], 3: [(18, 26)]} + for worker_id, intervals in expected.items(): + result = calc_slice_intervals( + [("/data/test.parquet", 35)], + worker_id=worker_id, + num_workers=4, + batch_size=4, + equalize_rank_steps=True, + world_size=2, + ) + self.assertEqual(result, [intervals]) + + def test_calc_slice_intervals_two_ranks_resume(self): + # remaining [(100, 500), (600, 1000)] = 800 rows, 400 per rank; rank 1's + # share [400, 800) spans the gap between the two remaining intervals + checkpoint_state = { + "/data/test.parquet:0": 99, + "/data/test.parquet:500": 599, + } + result = calc_slice_intervals( + [("/data/test.parquet", 1000)], + worker_id=2, + num_workers=4, + batch_size=128, + equalize_rank_steps=True, + checkpoint_state=checkpoint_state, + world_size=2, + ) + # rank 1, worker 0 gets two full batches -> logical [400, 656) + self.assertEqual(result, [[(600, 856)]]) + def test_calc_remaining_intervals_no_checkpoint(self): """Test remaining intervals when no checkpoint exists.""" result = calc_remaining_intervals( @@ -244,7 +305,7 @@ def test_calc_slice_intervals_single_worker(self): num_workers=1, batch_size=1, checkpoint_state=checkpoint_state, - )["/data/test.parquet"] + )[0] self.assertEqual(result, [(100, 500), (600, 1000)]) def test_calc_slice_intervals_two_workers(self): @@ -263,7 +324,7 @@ def test_calc_slice_intervals_two_workers(self): num_workers=2, batch_size=1, checkpoint_state=checkpoint_state, - )["/data/test.parquet"] + )[0] # Worker 1 gets second half result_w1 = calc_slice_intervals( [("/data/test.parquet", 1000)], @@ -271,7 +332,7 @@ def test_calc_slice_intervals_two_workers(self): num_workers=2, batch_size=1, checkpoint_state=checkpoint_state, - )["/data/test.parquet"] + )[0] # Combined should cover all intervals total_rows_w0 = sum(end - start for start, end in result_w0) @@ -288,7 +349,7 @@ def test_calc_slice_intervals_empty_intervals(self): num_workers=2, batch_size=1, checkpoint_state=checkpoint_state, - )["/data/test.parquet"] + )[0] self.assertEqual(result, []) def test_calc_slice_intervals_topology_change(self): @@ -309,7 +370,7 @@ def test_calc_slice_intervals_topology_change(self): num_workers=3, batch_size=1, checkpoint_state=checkpoint_state, - )["/data/test.parquet"] + )[0] for start, end in result: total_rows += end - start From b1bad3c0de72ee0506b3d8480f5bbfa5e20c5309 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Sun, 20 Sep 2026 11:52:57 +0800 Subject: [PATCH 5/8] [chore] bump version to 1.4.12 Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01163rPXLGHiM74LxdbWjMox --- tzrec/version.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tzrec/version.py b/tzrec/version.py index de29ce9b9..518f5c3ed 100644 --- a/tzrec/version.py +++ b/tzrec/version.py @@ -9,4 +9,4 @@ # See the License for the specific language governing permissions and # limitations under the License. -__version__ = "1.4.10" +__version__ = "1.4.12" From 7a7e9ce0e6d29cd61596fa38de05c698eeaebc03 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Sun, 20 Sep 2026 15:03:18 +0800 Subject: [PATCH 6/8] [bugfix] fetch odps session counts per rank and tighten planner docs and tests Rank 0 fetched every session's record count inside the pre-broadcast critical section, so a slow or failing readiness poll left the other ranks blocked in the collective; now the session ids are broadcast first and every rank fetches its own counts outside any collective. The planner lays the topped-up worker's block right after its carry chunk so it reads one interval per source, and calc_slice_intervals hands each source only its own checkpoint entries instead of rescanning the whole state per source. The single-row warning is replaced by a FAQ entry because it also fired in eval, where the knob does not apply; docs and docstrings are corrected and the tests assert disjoint tiling, the raw worker residue, and drop_remainder on its own. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01163rPXLGHiM74LxdbWjMox --- docs/source/faq.md | 16 +++++ docs/source/feature/data.md | 10 ++-- tzrec/datasets/dataset.py | 3 +- tzrec/datasets/kafka_dataset.py | 4 +- tzrec/datasets/odps_dataset.py | 25 ++++---- tzrec/datasets/parquet_dataset_test.py | 19 +++--- tzrec/datasets/utils.py | 27 +++++---- tzrec/datasets/utils_test.py | 83 +++++++++++++------------- 8 files changed, 106 insertions(+), 81 deletions(-) diff --git a/docs/source/faq.md b/docs/source/faq.md index 73bc73b47..6eaeb1138 100644 --- a/docs/source/faq.md +++ b/docs/source/faq.md @@ -546,3 +546,19 @@ ______________________________________________________________________ - TorchEasyRec在需要稠密序列时会自己按batch内的最大长度padding,并保留每条样本真实的长度用于mask,在tokenizer里补齐反而会丢掉这个信息 如果确实需要定长输出,注意TorchEasyRec生成的FG配置中`output_type`固定为`word_id`,因此只有`pad_id`生效:`pad_token`不会和`pad_id`做一致性校验,配错了不会报错;`pad_id`也不会校验是否在词表范围内,超出词表大小时训练会在embedding查表时越界。 + +**Q21: 训练报 ValueError: Expected more than 1 value per channel when training** + +**报错信息:** + +``` + File ".../torch/nn/functional.py", line ..., in batch_norm + _verify_batch_size(input.size()) + File ".../torch/nn/functional.py", line ..., in _verify_batch_size + raise ValueError( +ValueError: Expected more than 1 value per channel when training, got input size torch.Size([1, 1024]) +``` + +**原因:** 样本表的行数恰好使得某一轮训练的最后一个batch只有1行。BatchNorm在训练模式下需要在batch内计算统计量,1行样本无法计算方差;batch内负采样、listwise loss等也要求batch内至少有2行样本。TorchEasyRec默认不丢弃任何样本,因此最后一个不足batch_size的batch会原样送入模型。 + +**解决方法:** 在`data_config`中设置`min_batch_size: 2`,训练时行数小于2的最后一个batch会被丢弃,每个`proc`每轮最多丢弃1行样本;也可以设置`drop_remainder: true`丢弃所有不足batch_size的batch。两个参数都只在训练时生效,评估和预测不受影响。 diff --git a/docs/source/feature/data.md b/docs/source/feature/data.md index 3b901be21..5e94b440c 100644 --- a/docs/source/feature/data.md +++ b/docs/source/feature/data.md @@ -33,8 +33,6 @@ data_config { - `odps://{project}/tables/{table_name}/{partition}`,多表按逗号分隔 - 如果单表需要设置多个分区,可以用`&`简写,来分隔多个分区,`odps://{project}/tables/{table_name}/{partition1}&{partition2}` -- 注意: 训练和评估时,OdpsDataset会保证每个`proc`读取相同的步数以避免同步训练卡住,为此全部输入合计最多有`nproc - 1`行样本不会被读取;预测时不受影响 - - 运行训练/评估/导出/预测等命令时 - **本地环境**: @@ -77,8 +75,6 @@ data_config { - 注意: 如果每个parquet文件中的数据量不相等或文件数据小于worker数,ParquetDataset会自动重分配数据,来保证每个worker读取的数据量相等。但仍建议parquet文件数是 `nproc-per-node * nnodes * num_workers`的倍数,并且每个parquet文件的数据量基本相等,减少数据自动重分配的IO开销。 -- 注意: 训练和评估时,ParquetDataset会保证每个`proc`读取相同的步数以避免同步训练卡住,为此全部输入合计最多有`nproc - 1`行样本不会被读取;预测时不受影响 - ## CsvDataset - input_path: 按如下格式设置 @@ -334,11 +330,13 @@ pipeline.global-job-parameters: | ### drop_remainder - 是否丢弃掉最后一个不足batch_size的batch数据,默认为false -- 仅在训练时生效,评估和预测时不会丢弃任何数据;等价于`min_batch_size`设置为`batch_size` +- 仅在训练时生效;等价于`min_batch_size`设置为`batch_size`;评估和预测不会因最后一个batch不足而丢弃数据 ### min_batch_size -- 训练时,丢弃掉每个数据读取进程最后一个行数小于`min_batch_size`的batch,默认为0(不丢弃);评估和预测时不生效 +- 训练时,丢弃掉最后一个行数小于`min_batch_size`的batch,默认为0(不丢弃);评估和预测时不生效 +- OdpsDataset和ParquetDataset按`proc`(rank)规划数据,每个`proc`每轮最多丢弃`min_batch_size - 1`行;CsvDataset按数据读取进程丢弃 +- 须满足`min_batch_size <= batch_size`,否则训练时报错 - 使用BatchNorm等要求batch内至少2行样本的模型时,建议设置为2,避免样本表行数恰好使最后一个batch只剩1行导致训练失败 ### batch_cost_size diff --git a/tzrec/datasets/dataset.py b/tzrec/datasets/dataset.py index 028b4d23d..8bcd3a4e3 100644 --- a/tzrec/datasets/dataset.py +++ b/tzrec/datasets/dataset.py @@ -552,7 +552,8 @@ class BaseReader(metaclass=_reader_meta_cls): batch_cost_size (int): batch cost limit size. min_batch_size (int): drop a final batch with fewer rows, 0 disables. equalize_rank_steps (bool): make every rank yield the same number of - batches, honored by readers that slice rows by count. + batches; honored by readers that slice rows by count and not + guaranteed when batch_cost_size cuts batches by cost. """ def __init__( diff --git a/tzrec/datasets/kafka_dataset.py b/tzrec/datasets/kafka_dataset.py index 1a614a056..cea61665a 100644 --- a/tzrec/datasets/kafka_dataset.py +++ b/tzrec/datasets/kafka_dataset.py @@ -257,7 +257,8 @@ class KafkaReader(BaseReader): input_path (str): kafka URI, e.g. kafka://broker:9092/topic?group.id=xxx batch_size (int): batch size. selected_cols (list): selection column names. - drop_remainder (bool): drop last batch. + drop_remainder (bool): inert, the consume loop never ends so no tail + batch is formed; min_batch_size is inert for the same reason. input_fields (list): list of pa.Field for schema definition. """ @@ -376,7 +377,6 @@ def _reader( Args: worker_id: Worker ID num_workers: Total number of workers - world_size: Unused, partitions are split by global worker id Yields: PyArrow RecordBatch diff --git a/tzrec/datasets/odps_dataset.py b/tzrec/datasets/odps_dataset.py index 22ccaf9dd..294f2aa3b 100644 --- a/tzrec/datasets/odps_dataset.py +++ b/tzrec/datasets/odps_dataset.py @@ -522,22 +522,19 @@ def _init_session(self) -> None: else: session_ids.append(None) - # a session's record count is an immutable snapshot: fetch it once - # here instead of in every dataloader worker - sess_infos = [(x, None) for x in session_ids] - if int(os.environ.get("RANK", 0)) == 0: - sess_reqs = [SessionRequest(session_id=x) for x in session_ids] - with ThreadPoolExecutor(max_workers=8) as executor: - record_counts = executor.map( + if self._pg is not None: + dist.broadcast_object_list(session_ids, group=self._pg) + sess_reqs = [SessionRequest(session_id=x) for x in session_ids] + self._input_to_sess[input_path] = sess_reqs + # a session's record count is an immutable snapshot: every rank fetches + # it once here, outside any collective, instead of in every worker + with ThreadPoolExecutor(max_workers=8) as executor: + record_counts = list( + executor.map( _get_session_record_count, [client] * len(sess_reqs), sess_reqs ) - sess_infos = list(zip(session_ids, record_counts)) - if self._pg is not None: - dist.broadcast_object_list(sess_infos, group=self._pg) - self._input_to_sess[input_path] = [ - SessionRequest(session_id=x) for x, _ in sess_infos - ] - self._sess_row_counts.update(dict(sess_infos)) + ) + self._sess_row_counts.update(zip(session_ids, record_counts)) # refresh session if int(os.environ.get("RANK", 0)) == 0: t = threading.Thread( diff --git a/tzrec/datasets/parquet_dataset_test.py b/tzrec/datasets/parquet_dataset_test.py index 4daabc5af..3d1d2605a 100644 --- a/tzrec/datasets/parquet_dataset_test.py +++ b/tzrec/datasets/parquet_dataset_test.py @@ -316,14 +316,17 @@ def _drain(iterator): @parameterized.expand( [ # 8200 rows over 8 workers: one 8-row tail instead of eight 1-row tails - [0, [1024] * 8 + [8]], - [2, [1024] * 8 + [8]], - # drop_remainder drops the 8-row tail only - [1024, [1024] * 8], + [0, False, [1024] * 8 + [8]], + [2, False, [1024] * 8 + [8]], + # drop_remainder alone, and min_batch_size=batch_size, drop the tail only + [0, True, [1024] * 8], + [1024, False, [1024] * 8], ], name_func=parameterized_name_func, ) - def test_create_dataloader_tail_batches(self, min_batch_size, expected): + def test_create_dataloader_tail_batches( + self, min_batch_size, drop_remainder, expected + ): feature_cfgs = self._create_feature_cfgs() features = create_features(feature_cfgs) with tempfile.TemporaryDirectory(prefix="tzrec_") as test_dir: @@ -335,7 +338,7 @@ def test_create_dataloader_tail_batches(self, min_batch_size, expected): label_fields=["label"], num_workers=8, min_batch_size=min_batch_size, - drop_remainder=min_batch_size == 1024, + drop_remainder=drop_remainder, ) dataloader = create_dataloader(data_config, features, f"{test_dir}/*") sizes = sorted( @@ -347,7 +350,9 @@ def test_create_dataloader_tail_batches(self, min_batch_size, expected): @parameterized.expand( [[Mode.EVAL], [Mode.PREDICT]], name_func=parameterized_name_func ) - def test_create_dataloader_eval_predict_keep_every_row(self, mode): + def test_create_dataloader_eval_predict_no_tail_drop(self, mode): + # single rank: the knobs are inert outside training; multi-rank eval may + # still skip rows for lockstep, see test_parquet_reader_equal_steps feature_cfgs = self._create_feature_cfgs() features = create_features(feature_cfgs) with tempfile.TemporaryDirectory(prefix="tzrec_") as test_dir: diff --git a/tzrec/datasets/utils.py b/tzrec/datasets/utils.py index a451a3a9f..9ad80a390 100644 --- a/tzrec/datasets/utils.py +++ b/tzrec/datasets/utils.py @@ -906,21 +906,27 @@ def _rows_before(num_rows: int, r: int) -> int: for start, rows in shares: chunks: List[Tuple[int, int, int]] = [] pos = 0 + topped = None if carry and rows: pos = min(batch_size - carry, rows) chunks.append((holder, 0, pos)) loads[holder] += pos carry = (carry + pos) % batch_size + topped = holder full, tail = divmod(rows - pos, batch_size) q, r = divmod(full, num_workers) order = sorted(range(num_workers), key=lambda w: (loads[w], w)) counts = {w: (q + 1 if i < r else q) * batch_size for i, w in enumerate(order)} for w in order: loads[w] += counts[w] - # the tail opens the next batch on the least loaded worker, whose block - # is laid out last so that block and tail form one interval + # the topped-up worker's block follows its carry chunk and the tail opens + # the next batch on the least loaded worker, whose block is laid out + # last, so that each of them reads one interval from this source owner = min(range(num_workers), key=lambda w: (loads[w], w)) - for w in order: + layout = [w for w in order if w != topped] + if topped is not None: + layout.insert(0, topped) + for w in layout: if counts[w] and (w != owner or not tail): chunks.append((w, pos, pos + counts[w])) pos += counts[w] @@ -1006,8 +1012,14 @@ def calc_slice_intervals( local_workers = num_workers // world_size rank, local_worker_id = divmod(worker_id, local_workers) + # hand every source only its own checkpoint entries, in one pass over the keys + state_by_prefix: Dict[str, Dict[str, int]] = {} + for key, consumed in (checkpoint_state or {}).items(): + prefix, sep, _ = key.rpartition(":") + if sep: + state_by_prefix.setdefault(prefix, {})[key] = consumed remaining = [ - calc_remaining_intervals(checkpoint_state, prefix, total_rows) + calc_remaining_intervals(state_by_prefix.get(prefix), prefix, total_rows) for prefix, total_rows in sources ] plan = plan_rank_worker_intervals( @@ -1028,13 +1040,6 @@ def calc_slice_intervals( ] for intervals, logical in zip(remaining, plan) ] - num_rows = sum(end - start for intervals in result for start, end in intervals) - if equalize_rank_steps and min_batch_size < 2 and num_rows % batch_size == 1: - logger.warning( - "The final training batch of this pass has a single row; set " - "data_config.min_batch_size >= 2 if the model needs more than one " - "row per batch, e.g. with BatchNorm." - ) return result diff --git a/tzrec/datasets/utils_test.py b/tzrec/datasets/utils_test.py index d51a3a0ed..2db8e5720 100644 --- a/tzrec/datasets/utils_test.py +++ b/tzrec/datasets/utils_test.py @@ -73,6 +73,7 @@ def _rank_batches( total += end - start num_read += total worker_totals.append(total) + assert not 0 < total % batch_size < min_batch_size, total # a worker's stream is whole batches plus at most one final tail batches.extend([batch_size] * (total // batch_size)) if total % batch_size >= max(min_batch_size, 1): @@ -106,8 +107,6 @@ def test_plan_rank_worker_intervals_invariants(self): f"bs={batch_size} min_bs={min_batch_size}" ) for batches in per_rank: - self.assertTrue(all(b <= batch_size for b in batches), msg) - self.assertTrue(all(b >= min_batch_size for b in batches), msg) # at most one partial batch per rank and pass self.assertLessEqual( sum(1 for b in batches if b != batch_size), 1, msg @@ -308,36 +307,47 @@ def test_calc_slice_intervals_single_worker(self): )[0] self.assertEqual(result, [(100, 500), (600, 1000)]) + def _assert_tiles(self, slices, remaining): + """Slices are pairwise disjoint and their union is the remaining rows.""" + rows = [r for s in slices for start, end in s for r in range(start, end)] + self.assertEqual(len(rows), len(set(rows))) + self.assertEqual( + sorted(rows), [r for start, end in remaining for r in range(start, end)] + ) + def test_calc_slice_intervals_two_workers(self): - """Test calc_slice_intervals among two workers.""" - # Total remaining: 800 rows (400 + 400) - # Intervals: [(100, 500), (600, 1000)] + """Two even shares tile the remaining intervals without overlap.""" checkpoint_state = { "/data/test.parquet:0": 99, "/data/test.parquet:500": 599, } - - # Worker 0 gets first half of total rows - result_w0 = calc_slice_intervals( - [("/data/test.parquet", 1000)], - worker_id=0, - num_workers=2, - batch_size=1, - checkpoint_state=checkpoint_state, - )[0] - # Worker 1 gets second half - result_w1 = calc_slice_intervals( - [("/data/test.parquet", 1000)], - worker_id=1, - num_workers=2, - batch_size=1, - checkpoint_state=checkpoint_state, - )[0] - - # Combined should cover all intervals - total_rows_w0 = sum(end - start for start, end in result_w0) - total_rows_w1 = sum(end - start for start, end in result_w1) - self.assertEqual(total_rows_w0 + total_rows_w1, 800) + slices = [ + calc_slice_intervals( + [("/data/test.parquet", 1000)], + worker_id=worker_id, + num_workers=2, + batch_size=1, + checkpoint_state=checkpoint_state, + )[0] + for worker_id in range(2) + ] + self._assert_tiles(slices, [(100, 500), (600, 1000)]) + + def test_calc_slice_intervals_two_sources_resume(self): + """Even shares over two sources, the first partly consumed.""" + sources = [("a", 1000), ("b", 500)] + slices = [ + calc_slice_intervals( + sources, + worker_id=worker_id, + num_workers=2, + batch_size=128, + checkpoint_state={"a:0": 399}, + ) + for worker_id in range(2) + ] + self._assert_tiles([s[0] for s in slices], [(400, 1000)]) + self._assert_tiles([s[1] for s in slices], [(0, 500)]) def test_calc_slice_intervals_empty_intervals(self): """Test calc_slice_intervals with empty intervals (fully consumed).""" @@ -353,30 +363,23 @@ def test_calc_slice_intervals_empty_intervals(self): self.assertEqual(result, []) def test_calc_slice_intervals_topology_change(self): - """Test calc_slice_intervals when changing from 2 to 3 workers.""" - # Original 2 workers, remaining intervals from their checkpoints - # Intervals: [(300, 500), (800, 1000)] = 400 rows total + """Resuming with 3 workers tiles what 2 workers left behind.""" checkpoint_state = { "/data/test.parquet:0": 299, "/data/test.parquet:500": 799, } - - # Now redistribute among 3 workers - total_rows = 0 - for worker_id in range(3): - result = calc_slice_intervals( + slices = [ + calc_slice_intervals( [("/data/test.parquet", 1000)], worker_id=worker_id, num_workers=3, batch_size=1, checkpoint_state=checkpoint_state, )[0] - for start, end in result: - total_rows += end - start - - self.assertEqual(total_rows, 400) # All remaining rows accounted for + for worker_id in range(3) + ] + self._assert_tiles(slices, [(300, 500), (800, 1000)]) - # Every case verifies output, input non-mutation, and dict identity. @parameterized.expand( [ # (name, input_data, item_id_field, user_id_field, From 0373c9dc54b92b60cd1ed2fda91ab2939e7797a2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Sun, 20 Sep 2026 15:21:01 +0800 Subject: [PATCH 7/8] [refactor] read one contiguous chunk per worker and source A worker's buffer only depends on how many rows it takes from a source, not on where they sit, so the planner now gives every worker a single chunk per source: its whole batches, plus the rows completing the open batch for the holder, plus the tail for the least loaded worker. This drops the layout ordering and adjacency merge that previously left the holder with two intervals when it also owned the tail. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01163rPXLGHiM74LxdbWjMox --- tzrec/datasets/utils.py | 64 +++++++++++++----------------------- tzrec/datasets/utils_test.py | 4 +-- 2 files changed, 25 insertions(+), 43 deletions(-) diff --git a/tzrec/datasets/utils.py b/tzrec/datasets/utils.py index 9ad80a390..2c5fe8468 100644 --- a/tzrec/datasets/utils.py +++ b/tzrec/datasets/utils.py @@ -844,10 +844,10 @@ def plan_rank_worker_intervals( ``min_batch_size`` is cut off the end of the stream, on every rank alike. 3. Inside a rank the stream is cut into whole batches plus one tail and dealt to the workers source by source: the rows completing the open batch go to - its holder, whole batches are dealt as contiguous blocks to the least - loaded workers first, and the tail opens the next batch on the least - loaded worker. Only the holder ever buffers a partial batch, so a rank - ends a pass with at most one. + its holder, whole batches go to the least loaded workers first, and the + tail opens the next batch on the least loaded worker. Every worker reads + one contiguous chunk per source, and only the holder ever buffers a + partial batch, so a rank ends a pass with at most one. Dropped rows are never read. @@ -862,8 +862,8 @@ def plan_rank_worker_intervals( min_batch_size (int): drop a final batch with fewer rows, 0 disables. Returns: - for every source, the (start, end) intervals this worker reads, in read - order; empty when the worker reads nothing from that source. + for every source, a list holding the (start, end) interval this worker + reads, empty when it reads nothing from that source. """ assert 0 < num_workers and 0 <= worker_id < num_workers assert 0 < world_size and 0 <= rank < world_size @@ -904,46 +904,28 @@ def _rows_before(num_rows: int, r: int) -> int: loads = [0] * num_workers holder, carry = 0, 0 for start, rows in shares: - chunks: List[Tuple[int, int, int]] = [] - pos = 0 - topped = None - if carry and rows: - pos = min(batch_size - carry, rows) - chunks.append((holder, 0, pos)) - loads[holder] += pos - carry = (carry + pos) % batch_size - topped = holder - full, tail = divmod(rows - pos, batch_size) + # a worker's buffer only cares how many rows it gets from a source, so + # every worker reads one contiguous chunk: its whole batches, plus the + # rows completing the open batch for its holder, plus the tail for the + # least loaded worker, which then holds the next open batch + need = min(batch_size - carry, rows) if carry else 0 + loads[holder] += need + carry = (carry + need) % batch_size + full, tail = divmod(rows - need, batch_size) q, r = divmod(full, num_workers) order = sorted(range(num_workers), key=lambda w: (loads[w], w)) - counts = {w: (q + 1 if i < r else q) * batch_size for i, w in enumerate(order)} - for w in order: - loads[w] += counts[w] - # the topped-up worker's block follows its carry chunk and the tail opens - # the next batch on the least loaded worker, whose block is laid out - # last, so that each of them reads one interval from this source - owner = min(range(num_workers), key=lambda w: (loads[w], w)) - layout = [w for w in order if w != topped] - if topped is not None: - layout.insert(0, topped) - for w in layout: - if counts[w] and (w != owner or not tail): - chunks.append((w, pos, pos + counts[w])) - pos += counts[w] + chunk = [0] * num_workers + for i, w in enumerate(order): + chunk[w] = (q + 1 if i < r else q) * batch_size + loads[w] += chunk[w] + chunk[holder] += need if tail: - chunks.append((owner, pos, pos + counts[owner] + tail)) + owner = min(range(num_workers), key=lambda w: (loads[w], w)) + chunk[owner] += tail loads[owner] += tail holder, carry = owner, tail - - mine: List[Tuple[int, int]] = [] - for w, lo, hi in chunks: - if w != worker_id: - continue - if mine and mine[-1][1] == start + lo: - mine[-1] = (mine[-1][0], start + hi) - else: - mine.append((start + lo, start + hi)) - result.append(mine) + pos = start + sum(chunk[:worker_id]) + result.append([(pos, pos + chunk[worker_id])] if chunk[worker_id] else []) return result diff --git a/tzrec/datasets/utils_test.py b/tzrec/datasets/utils_test.py index 2db8e5720..b5d730e7e 100644 --- a/tzrec/datasets/utils_test.py +++ b/tzrec/datasets/utils_test.py @@ -187,8 +187,8 @@ def test_plan_rank_worker_intervals_predict_keeps_rows(self): def test_calc_slice_intervals_two_ranks_two_workers(self): # 35 rows: rank 0 owns [0, 18) with the extra row, rank 1 owns [18, 35); # inside each rank the 4 full batches are dealt two per worker and the - # tail opens on worker 0, whose block is laid out last - expected = {0: [(8, 18)], 1: [(0, 8)], 2: [(26, 35)], 3: [(18, 26)]} + # tail opens on worker 0; chunks are laid out in worker order + expected = {0: [(0, 10)], 1: [(10, 18)], 2: [(18, 27)], 3: [(27, 35)]} for worker_id, intervals in expected.items(): result = calc_slice_intervals( [("/data/test.parquet", 35)], From 1008de9193b057d7dbacd24f94e95736da93e4d6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A4=A9=E9=82=91?= Date: Sun, 20 Sep 2026 17:13:39 +0800 Subject: [PATCH 8/8] [bugfix] keep unstarted leading rows of a source when resuming calc_remaining_intervals inferred each keyed chunk's end from the next key and attributed the rows after the last key to it, but rows before the first key were treated as consumed. A worker with fewer batches in one source reaches the next source before its peers, so a checkpoint can hold a key for a later chunk of that source while an earlier chunk has no key yet; on resume those earlier rows were skipped. Rows before the first key can only lack a key because no batch ever read them, so they are now returned as remaining. The planner tests also assert that every worker gets at most one interval per source. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01163rPXLGHiM74LxdbWjMox --- tzrec/datasets/utils.py | 5 ++++- tzrec/datasets/utils_test.py | 36 ++++++++++++++++++++++++++++++++++++ 2 files changed, 40 insertions(+), 1 deletion(-) diff --git a/tzrec/datasets/utils.py b/tzrec/datasets/utils.py index 2c5fe8468..39a762938 100644 --- a/tzrec/datasets/utils.py +++ b/tzrec/datasets/utils.py @@ -800,8 +800,11 @@ def calc_remaining_intervals( # Sort by start to infer original ranges entries.sort(key=lambda x: x[0]) - # Calculate remaining intervals + # Calculate remaining intervals; rows before the first keyed range were + # never read, a consumed range always leaves its own key remaining = [] + if entries[0][0] > 0: + remaining.append((0, entries[0][0])) num_entries = len(entries) for i, (_, consumed) in enumerate(entries): # Infer the end of this worker's range diff --git a/tzrec/datasets/utils_test.py b/tzrec/datasets/utils_test.py index b5d730e7e..6aac09199 100644 --- a/tzrec/datasets/utils_test.py +++ b/tzrec/datasets/utils_test.py @@ -66,6 +66,8 @@ def _rank_batches( ) total = 0 for t, source_intervals in enumerate(intervals): + # one contiguous chunk per worker and source + assert len(source_intervals) <= 1, source_intervals for start, end in source_intervals: assert 0 <= start < end <= rows[t] assert seen[t].isdisjoint(range(start, end)) @@ -278,6 +280,16 @@ def test_calc_remaining_intervals_fully_consumed(self): # No remaining intervals self.assertEqual(result, []) + def test_calc_remaining_intervals_unstarted_leading_range(self): + """A range with no key yet was never read, even before the first key.""" + checkpoint_state = {"/data/test.parquet:4": 7} + result = calc_remaining_intervals( + checkpoint_state=checkpoint_state, + input_path="/data/test.parquet", + total_rows=8, + ) + self.assertEqual(result, [(0, 4)]) + def test_calc_remaining_intervals_unrelated_path(self): """Test remaining intervals when checkpoint is for different path.""" checkpoint_state = {"/data/other.parquet:0": 499} @@ -349,6 +361,30 @@ def test_calc_slice_intervals_two_sources_resume(self): self._assert_tiles([s[0] for s in slices], [(400, 1000)]) self._assert_tiles([s[1] for s in slices], [(0, 500)]) + def test_calc_slice_intervals_worker_ahead_in_next_source(self): + """Rows of a source chunk nobody has started yet survive a resume. + + Sources [12, 8], 2 workers, batch 4: worker 1 has one batch in source 0 + and reaches source 1 first, so the checkpoint keys source1:4 but not + source1:0. + """ + sources = [("source0", 12), ("source1", 8)] + checkpoint_state = {"source0:0": 7, "source0:8": 11, "source1:4": 7} + slices = [ + calc_slice_intervals( + sources, + worker_id=worker_id, + num_workers=2, + batch_size=4, + equalize_rank_steps=True, + checkpoint_state=checkpoint_state, + world_size=1, + ) + for worker_id in range(2) + ] + self._assert_tiles([s[0] for s in slices], []) + self._assert_tiles([s[1] for s in slices], [(0, 4)]) + def test_calc_slice_intervals_empty_intervals(self): """Test calc_slice_intervals with empty intervals (fully consumed).""" # All data consumed: checkpoint at row 999 (last row)