(
cfg: DictConfig,
tokenizer: Tokenizer,
device_batch_size: int,
)
| 291 | |
| 292 | |
| 293 | def build_streaming_dataset( |
| 294 | cfg: DictConfig, |
| 295 | tokenizer: Tokenizer, |
| 296 | device_batch_size: int, |
| 297 | ): |
| 298 | # build streams |
| 299 | streams_dict = cfg.dataset.get("streams", None) |
| 300 | streams = None |
| 301 | if streams_dict is not None: |
| 302 | streams = [] |
| 303 | for _, stream in streams_dict.items(): |
| 304 | streams.append( |
| 305 | Stream( |
| 306 | remote=stream.get("remote", None) or cfg.dataset.get("remote", None), |
| 307 | local=stream.get("local", None) or cfg.dataset.get("local", None), |
| 308 | split=stream.get("split", None) or cfg.dataset.get("split", None), |
| 309 | proportion=stream.get("proportion", None), |
| 310 | repeat=stream.get("repeat", None), |
| 311 | choose=stream.get("choose", None), |
| 312 | download_retry=stream.get("download_retry", None) or cfg.dataset.get("download_retry", 2), |
| 313 | download_timeout=stream.get("download_timeout", None) or cfg.dataset.get("download_timeout", 60), |
| 314 | validate_hash=stream.get("validate_hash", None) or cfg.dataset.get("validate_hash", None), |
| 315 | keep_zip=stream.get("keep_zip", None) or cfg.dataset.get("keep_zip", False), |
| 316 | ) |
| 317 | ) |
| 318 | |
| 319 | # build dataset potentially with streams |
| 320 | dataset = StreamingTextDataset( |
| 321 | tokenizer=tokenizer, |
| 322 | max_seq_len=cfg.dataset.max_seq_len, |
| 323 | streams=streams, |
| 324 | remote=cfg.dataset.get("remote", None), |
| 325 | local=cfg.dataset.get("local", None), |
| 326 | split=cfg.dataset.get("split", None), |
| 327 | download_retry=cfg.dataset.get("download_retry", 2), |
| 328 | download_timeout=cfg.dataset.get("download_timeout", 60), |
| 329 | validate_hash=cfg.dataset.get("validate_hash", None), |
| 330 | keep_zip=cfg.dataset.get("keep_zip", False), |
| 331 | epoch_size=cfg.dataset.get("epoch_size", None), |
| 332 | predownload=cfg.dataset.get("predownload", 100_000), |
| 333 | partition_algo=cfg.dataset.get("partition_algo", "orig"), |
| 334 | num_canonical_nodes=cfg.dataset.get("num_canonical_nodes", 128), |
| 335 | batch_size=device_batch_size, |
| 336 | shuffle=cfg.dataset.get("shuffle", False), |
| 337 | shuffle_algo=cfg.dataset.get("shuffle_algo", "py1s"), |
| 338 | shuffle_seed=cfg.dataset.get("shuffle_seed", 9176), |
| 339 | cache_limit=cfg.dataset.get("cache_limit", None), |
| 340 | ) |
| 341 | return dataset |
| 342 | |
| 343 | |
| 344 | def build_no_streaming_dataset( |
no test coverage detected