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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 16 additions & 0 deletions docs/source/faq.md
Original file line number Diff line number Diff line change
Expand Up @@ -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。两个参数都只在训练时生效,评估和预测不受影响。
8 changes: 8 additions & 0 deletions docs/source/feature/data.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
9 changes: 6 additions & 3 deletions tzrec/datasets/csv_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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__(
Expand All @@ -119,6 +121,7 @@ def __init__(
shuffle_buffer_size,
sample_cost_field=sample_cost_field,
batch_cost_size=batch_cost_size,
**kwargs,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nit: CsvReader is the only reader whose docstring wasn't updated — it still says drop_remainder (bool): drop last batch. (line 90) and doesn't mention min_batch_size, which CsvDataset now passes (line 74) through this newly added **kwargs. BaseReader, OdpsReader, and ParquetReader all got the new wording in this PR.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Already updated in 91e147d: the CsvReader docstring now describes drop_remainder as min_batch_size=batch_size and lists min_batch_size.

)
self._csv_fmt = ds.CsvFileFormat(
parse_options=pa.csv.ParseOptions(delimiter=delimiter),
Expand Down Expand Up @@ -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]
Expand Down
46 changes: 37 additions & 9 deletions tzrec/datasets/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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.
Expand All @@ -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)
Expand Down Expand Up @@ -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__(
Expand All @@ -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
Expand All @@ -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(
Expand Down Expand Up @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion tzrec/datasets/dataset_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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())

Expand Down
6 changes: 4 additions & 2 deletions tzrec/datasets/kafka_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
"""

Expand Down Expand Up @@ -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
Expand Down
Loading
Loading