bernini-diffusers-v2-demo / veomni /data /dynamic_batching.py
multimodalart's picture
multimodalart HF Staff
Bernini-Diffusers-v2 r2v demo
fed6c68 verified
Raw
History Blame Contribute Delete
12.6 kB
# Copyright 2025 Bytedance Ltd. and/or its affiliates
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import copy
import sys
import traceback
from collections import deque
from typing import Any, Callable, Dict, Generator, Iterator, Optional
from ..utils import logging
logger = logging.get_logger(__name__)
# TODO: add state dict for buffer to resume training.
class DynBszBuffer:
"""
A buffer to store samples for dynamic batch size.
"""
def __init__(self):
self._buffer = []
self._buffer_sample_lens = []
self.del_idxs = []
self.cur_idx = 0
self.all_token_cnt = 0
def append(self, item: Dict[str, Any]):
"""
Append a sample to the buffer.
Args:
item: a sample to append to the buffer.
The sample should be a dict containing an ``attention_mask`` tensor
whose ``.sum()`` gives the number of valid tokens for batching.
"""
self._buffer.append(item)
if "attention_mask" not in item:
raise KeyError("Expected 'attention_mask' in item")
self._buffer_sample_lens.append(item["attention_mask"].sum())
self.all_token_cnt += self._buffer_sample_lens[-1]
def get_samples(self, n_token_per_iter: int, force: bool = True):
"""
get samples from the buffer.
Args:
n_token_per_iter: the number of tokens to get.
force: if True, the first sample will be returned even if it is not full.
Returns:
samples: a list of samples.
"""
cum_seq_len = 0
samples = []
while self.cur_idx < len(self._buffer) and cum_seq_len < n_token_per_iter:
seq_len = self._buffer_sample_lens[self.cur_idx]
if self.cur_idx not in self.del_idxs and (
(force is True and cum_seq_len == 0) or (seq_len <= n_token_per_iter - cum_seq_len)
):
cum_seq_len += seq_len
samples.append(self._buffer[self.cur_idx])
self.del_idxs.append(self.cur_idx)
self.cur_idx += 1
assert len(samples) > 0
return samples
def __len__(self):
return len(self._buffer)
def flush(self):
""" "
Flush the buffer.
"""
self.cur_idx = 0
self.all_token_cnt -= sum([self._buffer_sample_lens[idx] for idx in self.del_idxs])
buffer_len = len(self._buffer)
self._buffer = [self._buffer[idx] for idx in range(buffer_len) if idx not in self.del_idxs]
self._buffer_sample_lens = [
self._buffer_sample_lens[idx] for idx in range(buffer_len) if idx not in self.del_idxs
]
self.del_idxs = []
def merge(self, buffer_to_merge: "DynBszBuffer"):
""" "
Merge the buffer with another buffer.
Args:
buffer_to_merge: the buffer to merge.
"""
self.flush()
buffer_to_merge.flush()
for item in buffer_to_merge._buffer:
self.append(item)
class BaseBatchingStrategy:
"""
Base class for batching strategy.
"""
def is_ready_for_micro_batch(self) -> bool:
raise NotImplementedError("should implement `is_ready_for_micro_batch`")
def put_item(self, item: Dict[str, Any]):
raise NotImplementedError("should implement `put_item`")
def get_micro_batch(self, step: int) -> Any:
raise NotImplementedError("should implement `get_micro_batch` ")
def empty(self) -> bool:
raise NotImplementedError("should implement `empty`")
class TextBatchingStrategy(BaseBatchingStrategy):
""" "
Batching strategy for text data.
Args:
token_micro_bsz: the number of tokens to get for each request.
bsz_warmup_steps: the number of steps to warm up the batch size.
bsz_warmup_init_mbtoken: the initial number of tokens to get for each request.
buffer_size: the size of the buffer.
"""
def __init__(
self,
token_micro_bsz,
buffer_size: int = 500,
bsz_warmup_steps: int = 0,
bsz_warmup_init_mbtoken: int = 200,
) -> None:
super().__init__()
self._step = 0
self.token_micro_bsz = token_micro_bsz
self.bsz_warmup_steps = bsz_warmup_steps
self.bsz_warmup_init_mbtoken = bsz_warmup_init_mbtoken
if bsz_warmup_steps > 0:
assert self.bsz_warmup_init_mbtoken > 0
self.buffer_size = buffer_size # minimum samples in buffer
self.buffer = DynBszBuffer()
def is_ready_for_micro_batch(self) -> bool:
return len(self.buffer) >= self.buffer_size and self.buffer.all_token_cnt >= self.token_micro_bsz
def put_item(self, item: Dict[str, Any]):
if item["input_ids"].shape[-1] <= 1:
print("WARNING: EMPTY STRING.")
return
self.buffer.append(item)
def get_cur_token_micro_bsz(self):
warmup = self.bsz_warmup_steps > 0 and self._step <= self.bsz_warmup_steps
if warmup:
return (
self.token_micro_bsz - self.bsz_warmup_init_mbtoken
) * self._step // self.bsz_warmup_steps + self.bsz_warmup_init_mbtoken
else:
return self.token_micro_bsz
def get_micro_batch(self, step) -> Any:
"""
Get a micro batch from the buffer according to the current step.
Args:
step: the current step.
Returns:
data: a list of samples.
"""
self._step = step
cur_token_micro_bsz = self.get_cur_token_micro_bsz()
samples = self.buffer.get_samples(cur_token_micro_bsz)
self.buffer.flush() # remove the selected samples.
return samples
def empty(self) -> bool:
return len(self.buffer) == 0
class DynamicBatchSizeDataLoader:
"""Dynamic batch DataLoader.
Args:
dataloader: torch DataLoader
batching_strategy: dynamic batch strategy
collate_fn: DataLoader collate_fn, collate data after get data from batching_strategy
num_micro_batch: num_micro_batch, if num_micro_batch == 1, return micro_batch for gradient accumulation
length: length of dataloader, if length == -1, length = sys.maxsize, default len(dataloader)
drop_last: if True, drop last batch if batch size < num_micro_batch
"""
def __init__(
self,
dataloader: Any,
batching_strategy: "BaseBatchingStrategy",
collate_fn: Optional[Callable] = None,
num_micro_batch: int = 1,
length: int = 0,
drop_last: bool = True,
) -> None:
self.batching_strategy = batching_strategy
self.num_micro_batch = num_micro_batch
self.dataloader_item_buffer = deque()
self.item_buffer = deque()
self.step = 0
self._collate_fn = collate_fn
self._dataloader = dataloader
self._drop_last = drop_last
self._data_iter: Iterator
self._resume = False
self._batch_data_iter: Generator
if length > 0:
self._length = length
elif length == -1:
self._length = sys.maxsize
else:
self._length = len(self._dataloader)
def __len__(self):
if self._length:
return self._length
else:
raise RuntimeError("length must set at init. before call len()")
def __iter__(self) -> Iterator:
if not self._resume:
self.step = 0
self._data_iter = iter(self._dataloader)
self._batch_data_iter = self.batch_data_generator()
self._resume = False
return self
def __next__(self):
return next(self._batch_data_iter)
def batch_data_generator(self):
batch = []
while True:
if self._length and self.step >= self._length:
return
if self.batching_strategy.is_ready_for_micro_batch():
micro_batch = self.batching_strategy.get_micro_batch(self.step)
if self._collate_fn:
micro_batch = self._collate_fn(micro_batch)
batch.append(micro_batch)
if len(batch) == self.num_micro_batch:
yield batch
self.step += 1
batch = []
try:
processing_item = next(self._data_iter)
except Exception as e:
if isinstance(e, StopIteration):
if self.step < self._length:
# call iter until reach length
self._data_iter = iter(self._dataloader)
processing_item = next(self._data_iter)
elif not self._drop_last and not self.batching_strategy.empty():
while not self.batching_strategy.empty():
micro_batch = self.batching_strategy.get_micro_batch(self.step)
if self._collate_fn:
micro_batch = self._collate_fn(micro_batch)
batch.append(micro_batch)
if len(batch) == self.num_micro_batch:
yield batch
self.step += 1
batch = []
while len(batch) < self.num_micro_batch:
padding_batch = copy.deepcopy(micro_batch)
padding_batch["padding_flag"] = True
batch.append(padding_batch)
yield batch
self.step += 1
return
else:
return
else:
logger.error(f"DynamicBatchDataset iter data exception: {e} \n{traceback.format_exc()}")
raise
# put processing_item to buffer
if isinstance(processing_item, dict):
processing_item = [processing_item]
for item in processing_item:
self.batching_strategy.put_item(item)
def state_dict(self):
# save state
state = self.__dict__.copy()
# remove internal fields
for k in list(state.keys()):
if k.startswith("_"):
del state[k]
# save dataloader state
if hasattr(self._dataloader, "state_dict"):
state["dataloader_state"] = self._dataloader.state_dict()
elif hasattr(self._dataloader, "__getstate__"):
state["dataloader_state"] = self._dataloader.__getstate__()
if hasattr(self.batching_strategy, "state_dict"):
state["batching_strategy_state"] = self.batching_strategy.state_dict() # type: ignore
del state["batching_strategy"]
return copy.deepcopy(state)
def load_state_dict(self, state: Dict[str, Any]):
if state["num_micro_batch"] != self.num_micro_batch:
logger.warning(
f"num_micro_batch changed: [ {state['num_micro_batch']} -> {self.num_micro_batch} ], will clear prefetch buffer"
)
del state["num_micro_batch"]
self.__dict__.update(state)
self._resume = True
if hasattr(self._dataloader, "load_state_dict"):
self._dataloader.load_state_dict(state["dataloader_state"])
elif hasattr(self._dataloader, "__getstate__"):
self._dataloader.__setstate__(state["dataloader_state"])
if "batching_strategy_state" in state:
self.batching_strategy.load_state_dict( # type: ignore
state["batching_strategy_state"]
)
del state["batching_strategy_state"]
self._data_iter = iter(self._dataloader)
self._batch_data_iter = self.batch_data_generator()
def set_epoch(self, epoch: int) -> None:
if hasattr(self._dataloader, "set_epoch"):
self._dataloader.set_epoch(epoch)