from abc import ABC, abstractmethod class BaseTask(ABC): def __init__(self, config): self.config = config self.dataset = None self.model = None self.trainer = None @abstractmethod def build_dataset(self): """构建并返回训练数据集实例""" pass @abstractmethod def build_model(self): """构建并返回模型实例""" pass def build_collator(self): """可选:构建数据整理器(collator)""" return None def build_sampler(self): """可选:构建采样器(支持动态批等)""" return None def build_trainer(self): """构建训练器(默认使用 HuggingFace Trainer)""" from transformers import Trainer trainer_args = self.config.get("trainer", {}).get("args", {}) return Trainer( model=self.model, args=trainer_args, train_dataset=self.dataset, data_collator=self.build_collator(), # **其余参数如 eval_dataset 可从 config 中补充 ) def run(self): """执行完整的训练流程:构建各组件并开始训练""" # 按序构建数据集、模型、collator、sampler 和训练器 self.dataset = self.build_dataset() self.model = self.build_model() self.collator = self.build_collator() self.sampler = self.build_sampler() self.trainer = self.build_trainer() # 开始训练 self.trainer.train()