diff --git a/docs/source/faq.md b/docs/source/faq.md index 73bc73b4..6eaeb113 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 74b1505e..5e94b440 100644 --- a/docs/source/feature/data.md +++ b/docs/source/feature/data.md @@ -330,6 +330,14 @@ pipeline.global-job-parameters: | ### drop_remainder - 是否丢弃掉最后一个不足batch_size的batch数据,默认为false +- 仅在训练时生效;等价于`min_batch_size`设置为`batch_size`;评估和预测不会因最后一个batch不足而丢弃数据 + +### min_batch_size + +- 训练时,丢弃掉最后一个行数小于`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/csv_dataset.py b/tzrec/datasets/csv_dataset.py index 4af8fe32..8ac99ff9 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, @@ -86,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__( @@ -119,6 +121,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), @@ -152,7 +155,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 11972bba..8bcd3a4e 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 + # 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: + raise ValueError( + 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 @@ -288,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 @@ -305,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. @@ -320,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) @@ -533,11 +545,15 @@ 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 number of + batches; honored by readers that slice rows by count and not + guaranteed when batch_cost_size cuts batches by cost. """ def __init__( @@ -550,12 +566,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 @@ -582,9 +602,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( @@ -624,7 +652,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/dataset_test.py b/tzrec/datasets/dataset_test.py index 91d3e280..c6ddb20e 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 c3557976..cea61665 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. """ @@ -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 fe594fc0..294f2aa3 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 @@ -387,13 +388,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 +406,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 number of batches. """ def __init__( @@ -426,7 +428,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,18 +442,19 @@ 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 = {} self._table_to_cli = {} self._input_to_sess = {} + self._sess_row_counts: Dict[str, int] = {} self._init_client() fields = [] @@ -522,9 +524,17 @@ def _init_session(self) -> None: if self._pg is not None: dist.broadcast_object_list(session_ids, group=self._pg) - self._input_to_sess[input_path] = [ - SessionRequest(session_id=x) for x in session_ids - ] + 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 + ) + ) + self._sess_row_counts.update(zip(session_ids, record_counts)) # refresh session if int(os.environ.get("RANK", 0)) == 0: t = threading.Thread( @@ -586,6 +596,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} " @@ -617,64 +628,49 @@ 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.""" - 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}", + self._sess_row_counts[sess_req.session_id], + client, + sess_req, ) + ) + plan = 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, + world_size=world_size, + ) - # 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), intervals in zip(sources, plan): + for start, end in intervals: + 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 461b0086..2ce50dca 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/odps_dataset_v1.py b/tzrec/datasets/odps_dataset_v1.py index 6ffd5c94..efeec946 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 69bbdc3e..8a1cef6e 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 number of batches. """ 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 @@ -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: @@ -248,18 +248,21 @@ 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, - ) + world_size=world_size, + )[0] 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 64b75620..3d1d2605 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,89 @@ 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, 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, drop_remainder, 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=drop_remainder, + ) + 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) + + @parameterized.expand( + [[Mode.EVAL], [Mode.PREDICT]], name_func=parameterized_name_func + ) + 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: + 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, + ) + num_rows = sum( + len(batch.reserves.get()) + if mode == Mode.PREDICT + else len(batch.labels["label"]) + 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}/*", mode=Mode.TRAIN) + def test_create_dataloader_get_iterator_reuses_eager_iterator(self): """get_iterator() yields the eagerly prefetched iterator first. @@ -436,6 +520,70 @@ def _reader_worker(rank, port): if p.exitcode != 0: raise RuntimeError(f"reader worker-{i} failed.") + @parameterized.expand( + [ + # 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, 1, [[4] * 4, [4] * 4]], + # 8200 rows over two ranks: 4100 each, 4-row tails + [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, local_workers, sizes + ): + def _reader_worker(rank, port): + 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, + ) + 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 + + 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() + procs = [ + mp.Process(target=_reader_worker, args=(rank, port)) for rank in range(2) + ] + for p in procs: + p.start() + for i, p in enumerate(procs): + p.join() + if p.exitcode != 0: + raise RuntimeError(f"reader worker-{i} failed.") + class ParquetWriterTest(unittest.TestCase): def setUp(self): diff --git a/tzrec/datasets/utils.py b/tzrec/datasets/utils.py index 6dc5cdef..39a76293 100644 --- a/tzrec/datasets/utils.py +++ b/tzrec/datasets/utils.py @@ -759,65 +759,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, @@ -859,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 @@ -876,79 +820,212 @@ def calc_remaining_intervals( return remaining if remaining else [] +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[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 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 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. + + 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 number of batches. + min_batch_size (int): drop a final batch with fewer rows, 0 disables. + + Returns: + 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 + assert 0 <= min_batch_size <= batch_size + + 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: + 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: + 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: + # 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)) + 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: + owner = min(range(num_workers), key=lambda w: (loads[w], w)) + chunk[owner] += tail + loads[owner] += tail + holder, carry = owner, tail + pos = start + sum(chunk[:worker_id]) + result.append([(pos, pos + chunk[worker_id])] if chunk[worker_id] else []) + 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. + world_size: Optional[int] = None, +) -> List[List[Tuple[int, int]]]: + """Assign the row intervals of every source to one of ``num_workers`` slices. - Flattens all intervals into a total row count, then assigns a portion - to each worker based on worker_id and num_workers. + 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: - 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): slice id. + num_workers (int): slice number. + batch_size (int): batch size. + 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. - input_path (str): the input path to filter checkpoint entries. + world_size (int, optional): number of ranks the slices are grouped into. 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. + for every source, in the order of ``sources``, the (start, end) intervals + of this slice. """ - 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, - ) - - 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 + if world_size is None: + rank, world_size, local_worker_id, local_workers = worker_id, num_workers, 0, 1 else: - result = [(worker_start, worker_end)] - - return result, total_remain + 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) + + # 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(state_by_prefix.get(prefix), 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, + ) + result = [ + [ + physical + for start, end in logical + for physical in _map_logical_range(intervals, start, end) + ] + for intervals, logical in zip(remaining, plan) + ] + 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 3e95a52a..6aac0919 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,202 @@ 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, 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): + intervals = 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, 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)) + seen[t].update(range(start, end)) + 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): + batches.append(total % batch_size) + per_rank.append(batches) + return per_rank, num_read, worker_totals + + 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, 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( + 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: + # 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: + # 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 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 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 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], + ], + 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_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; 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)], + 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.""" @@ -109,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} @@ -129,84 +310,112 @@ 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", - ) + )[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( - total_rows=1000, - worker_id=0, - num_workers=2, - checkpoint_state=checkpoint_state, - input_path="/data/test.parquet", - ) - # Worker 1 gets second half - result_w1, _ = calc_slice_intervals( - total_rows=1000, - worker_id=1, - num_workers=2, - checkpoint_state=checkpoint_state, - input_path="/data/test.parquet", - ) - - # 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_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) 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", - ) + )[0] 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( - total_rows=1000, + slices = [ + 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", - ) - for start, end in result: - total_rows += end - start - - self.assertEqual(total_rows, 400) # All remaining rows accounted for + )[0] + 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, diff --git a/tzrec/protos/data.proto b/tzrec/protos/data.proto index fb1af5d2..aa391336 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,6 +156,10 @@ message DataConfig { // whether dataloader batches are returned in first-in, first-out order optional bool in_order = 28 [default = true]; + // 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]; + // negative sampler oneof sampler { NegativeSampler negative_sampler = 101; diff --git a/tzrec/version.py b/tzrec/version.py index 1d5e245d..518f5c3e 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.11" +__version__ = "1.4.12"