msj19/gated_deltaproduct
05
1# -*- coding: utf-8 -*-2 3from __future__ import annotations4 5from dataclasses import dataclass, field6from typing import Optional7 8import transformers9from transformers import HfArgumentParser, TrainingArguments10 11from flame.logging import get_logger12 13logger = get_logger(__name__)14 15 16@dataclass17class TrainingArguments(TrainingArguments):18 19 model_name_or_path: str = field(20 default=None,21 metadata={22 "help": "Path to the model weight or identifier from huggingface.co/models or modelscope.cn/models."23 },24 )25 tokenizer: str = field(26 default="fla-hub/gla-1.3B-100B",27 metadata={"help": "Name of the tokenizer to use."}28 )29 use_fast_tokenizer: bool = field(30 default=False,31 metadata={"help": "Whether or not to use one of the fast tokenizer (backed by the tokenizers library)."},32 )33 from_config: bool = field(34 default=True,35 metadata={"help": "Whether to initialize models from scratch."},36 )37 dataset: Optional[str] = field(38 default=None,39 metadata={"help": "The dataset(s) to use. Use commas to separate multiple datasets."},40 )41 dataset_name: Optional[str] = field(42 default=None,43 metadata={"help": "The name of provided dataset(s) to use."},44 )45 cache_dir: str = field(46 default=None,47 metadata={"help": "Path to the cached tokenized dataset."},48 )49 split: str = field(50 default="train",51 metadata={"help": "Which dataset split to use for training and evaluation."},52 )53 streaming: bool = field(54 default=False,55 metadata={"help": "Enable dataset streaming."},56 )57 hf_hub_token: Optional[str] = field(58 default=None,59 metadata={"help": "Auth token to log in with Hugging Face Hub."},60 )61 preprocessing_num_workers: Optional[int] = field(62 default=None,63 metadata={"help": "The number of processes to use for the pre-processing."},64 )65 buffer_size: int = field(66 default=2048,67 metadata={"help": "Size of the buffer to randomly sample examples from in dataset streaming."},68 )69 context_length: int = field(70 default=2048,71 metadata={"help": "The context length of the tokenized inputs in the dataset."},72 )73 varlen: bool = field(74 default=False,75 metadata={"help": "Enable training with variable length inputs."},76 )77 78 79def get_train_args():80 parser = HfArgumentParser(TrainingArguments)81 args, unknown_args = parser.parse_args_into_dataclasses(return_remaining_strings=True)82 83 if unknown_args:84 print(parser.format_help())85 print("Got unknown args, potentially deprecated arguments: {}".format(unknown_args))86 raise ValueError("Some specified arguments are not used by the HfArgumentParser: {}".format(unknown_args))87 88 if args.should_log:89 transformers.utils.logging.set_verbosity(args.get_process_log_level())90 transformers.utils.logging.enable_default_handler()91 transformers.utils.logging.enable_explicit_format()92 # set seeds manually93 transformers.set_seed(args.seed)94 return args95 