(
cfg: DictConfig,
tokenizer: Tokenizer,
device_batch_size: int,
device_microbatch_size: int,
)
| 356 | |
| 357 | |
| 358 | def build_text_dataloader( |
| 359 | cfg: DictConfig, |
| 360 | tokenizer: Tokenizer, |
| 361 | device_batch_size: int, |
| 362 | device_microbatch_size: int, |
| 363 | ): |
| 364 | assert cfg.name == "text", f"Tried to build text dataloader with cfg.name={cfg.name}" |
| 365 | if cfg.dataset.get("group_method", None) is not None: |
| 366 | raise NotImplementedError( |
| 367 | "group_method is deprecated and has been removed.\nTo " |
| 368 | + "concatenate, use the --concat_tokens " |
| 369 | + "argument when creating your MDS dataset with convert_dataset.py" |
| 370 | ) |
| 371 | |
| 372 | if cfg.dataset.get("streaming", True): |
| 373 | dataset = build_streaming_dataset(cfg, tokenizer, device_batch_size) |
| 374 | sampler = None |
| 375 | else: |
| 376 | assert cfg.dataset.get("local", None) is not None, "Local path must be provided when not using streaming" |
| 377 | # sequence packing should never use padded sequences, regular dataloaders may if tokenizing on the fly |
| 378 | dataset = build_no_streaming_dataset( |
| 379 | cfg, tokenizer=tokenizer, pad_sequences=not cfg.get("sequence_packing", False) |
| 380 | ) |
| 381 | sampler = DistributedSamplerPCG64DXSM( |
| 382 | dataset, |
| 383 | num_replicas=dist.get_world_size(), |
| 384 | rank=dist.get_global_rank(), |
| 385 | shuffle=cfg.dataset.get("shuffle", False), |
| 386 | seed=cfg.dataset.get("shuffle_seed", 9176), |
| 387 | drop_last=cfg.drop_last, |
| 388 | ) |
| 389 | |
| 390 | mlm_probability = cfg.dataset.get("mlm_probability", None) |
| 391 | # only use sequence packing if using the no_streaming_dataset |
| 392 | if not cfg.dataset.get("streaming", True) and cfg.get("sequence_packing", False): |
| 393 | dataloader = DataLoader( |
| 394 | dataset, |
| 395 | collate_fn=lambda x: x, |
| 396 | batch_size=device_batch_size, |
| 397 | drop_last=False, |
| 398 | num_workers=cfg.num_workers, |
| 399 | pin_memory=cfg.get("pin_memory", True), |
| 400 | prefetch_factor=cfg.get("prefetch_factor", 2), |
| 401 | persistent_workers=cfg.get("persistent_workers", True), |
| 402 | timeout=cfg.get("timeout", 0), |
| 403 | sampler=sampler, |
| 404 | ) |
| 405 | sequence_packer = GreedyBestFitSequencePacker.from_composer( |
| 406 | dataloader, |
| 407 | batch_size=device_batch_size, |
| 408 | micro_batch_size=device_microbatch_size, |
| 409 | max_seq_len=cfg.dataset.max_seq_len, |
| 410 | buffer_size=cfg.get("packing_buffer_size", 5 * device_batch_size), |
| 411 | mask_token_id=tokenizer.mask_token_id, |
| 412 | pad_token_id=tokenizer.pad_token_id, |
| 413 | mask_prob=mlm_probability, |
| 414 | seed=cfg.dataset.get("shuffle_seed", 42), |
| 415 | batch_size_warmup_min_size=cfg.get("batch_size_warmup_min_size", None), |
no test coverage detected