feat(paralell): 添加分布式训练配置与并行工具支持
This commit is contained in:
@@ -9,7 +9,6 @@ if TYPE_CHECKING:
|
||||
|
||||
@dataclass
|
||||
class TrainConfig:
|
||||
|
||||
strategy: "BaseStrategy" = field(
|
||||
default=None,
|
||||
metadata={"help": "Training strategy."}
|
||||
@@ -54,6 +53,8 @@ class TrainConfig:
|
||||
default=1.0,
|
||||
metadata={"help": "Maximum gradient norm."}
|
||||
)
|
||||
|
||||
# dataloader setting
|
||||
random_seed: int = field(
|
||||
default=3407,
|
||||
metadata={"help": "Random seed."}
|
||||
@@ -69,4 +70,10 @@ class TrainConfig:
|
||||
pin_memory: bool = field(
|
||||
default=False,
|
||||
metadata={"help": "Pin memory for dataloader."}
|
||||
)
|
||||
|
||||
# distributed training
|
||||
nprocs: int = field(
|
||||
default=1,
|
||||
metadata={"help": "Number of processes for distributed training."}
|
||||
)
|
||||
Reference in New Issue
Block a user