MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / build_text_dataloader

Function build_text_dataloader

src/text_data.py:358–444  ·  view source on GitHub ↗
(
    cfg: DictConfig,
    tokenizer: Tokenizer,
    device_batch_size: int,
    device_microbatch_size: int,
)

Source from the content-addressed store, hash-verified

356
357
358def 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),

Callers 1

text_data.pyFile · 0.85

Tested by

no test coverage detected