CoolFace
Modelpublic

msj19/gated_deltaproduct

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes5downloads
parser.py95 linesDownload Raw Back to flame
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