INNER CODE UNIT · Python

minimum_n_tokens

Guitaricet/relora · torchrun_main.py:447

        minimum_n_tokens = args.total_batch_size * args.num_training_steps
        dataset_n_tokens = len(train_dataset) * args.max_length
        if dataset_n_tokens < minimum_n_tokens:
            raise ValueError(f"Dataset only has {dataset_n_tokens} tokens, but we need at least {minimum_n_tokens}")

        logger.info("Loading dataset preprocessing args to check on seq_length")
        with open(os.path.join(args.dataset_path, "args.json")) as f:
            dataset_preprocessing_args = json.load(f)
        assert dataset_preprocessing_args["sequence_length"] == args.max_length
        logger.info("All good! Loading tokenizer now")
        # ##############################
        tokenizer = AutoTokenizer.from_pretrained(
            dataset_preprocessing_args["tokenizer"],
            model_max_length=args.max_length,
        )
        logger.info("Tokenizer loaded")

    elif args.megatron_dataset_config is not None:

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…