mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
54 lines
1.3 KiB
Python
54 lines
1.3 KiB
Python
import dataclasses
|
|
import os
|
|
import sys
|
|
from dataclasses import dataclass, field
|
|
from enum import Enum, auto
|
|
from typing import Any, Dict, List, Literal, NewType, Optional, Tuple
|
|
|
|
from transformers import MODEL_FOR_CAUSAL_LM_MAPPING, HfArgumentParser
|
|
|
|
|
|
class ExperimentSetup(Enum):
|
|
LORA = "lora"
|
|
HYPER_LORA = "hyper_lora"
|
|
FULL_FINETUNE = "full_finetune"
|
|
|
|
|
|
@dataclass
|
|
class ModelArguments:
|
|
"""
|
|
Arguments for the base model.
|
|
"""
|
|
|
|
model_name_or_path: Optional[str] = field(
|
|
default=None,
|
|
metadata={"help": ("Base model name or path.")},
|
|
)
|
|
# use_peft: bool = field(
|
|
# default=False,
|
|
# metadata={"help": ("Whether to use PEFT or not for training.")},
|
|
# )
|
|
|
|
|
|
@dataclass
|
|
class LoRAArguments:
|
|
r: Optional[int] = field(
|
|
default=8,
|
|
metadata={"help": ("LoRA R value.")},
|
|
)
|
|
lora_dropout: Optional[float] = field(
|
|
default=0.05,
|
|
metadata={"help": ("LoRA dropout.")},
|
|
)
|
|
target_modules: Optional[list[str]] = field(
|
|
default=None,
|
|
metadata={"help": ("LoRA target modules.")},
|
|
)
|
|
|
|
|
|
@dataclass
|
|
class CtxTrainingArguments:
|
|
exp_setup: ExperimentSetup = field(
|
|
default=ExperimentSetup.LORA,
|
|
metadata={"help": "Experiment setup - LoRA, HyperLoRA, or full finetuning"},
|
|
)
|