Initial commit
This commit is contained in:
@@ -0,0 +1,11 @@
|
||||
from khaosz.trainer.dataset import DatasetLoader
|
||||
from khaosz.trainer.trainer import Trainer
|
||||
from khaosz.trainer.strategy import TrainConfig, CosineScheduleConfig, SgdrScheduleConfig
|
||||
|
||||
__all__ = [
|
||||
"DatasetLoader",
|
||||
"Trainer",
|
||||
"TrainConfig",
|
||||
"CosineScheduleConfig",
|
||||
"SgdrScheduleConfig",
|
||||
]
|
||||
@@ -0,0 +1,210 @@
|
||||
import torch
|
||||
import bisect
|
||||
import pickle as pkl
|
||||
from abc import ABC, abstractmethod
|
||||
from torch import Tensor
|
||||
from torch.utils.data import Dataset
|
||||
from typing import Callable, List, Dict, Literal, Union
|
||||
|
||||
MutiSeg = Dict[str, List[Tensor]]
|
||||
Seg = Dict[str, Tensor]
|
||||
|
||||
def load_pkl_files(paths: List[str]):
|
||||
segments: MutiSeg = {}
|
||||
total_samples = 0
|
||||
|
||||
for path in paths:
|
||||
with open(path, "rb") as f:
|
||||
pkl_file: Seg = pkl.load(f)
|
||||
for key, value in pkl_file.items():
|
||||
if key not in segments:
|
||||
segments[key] = []
|
||||
segments[key].append(value)
|
||||
first_key = list(pkl_file.keys())[0]
|
||||
total_samples += pkl_file[first_key].numel()
|
||||
|
||||
return segments, total_samples
|
||||
|
||||
|
||||
class BaseSegmentFetcher:
|
||||
def __init__(self, segments: List[Tensor]):
|
||||
self.segments = segments
|
||||
self.cum_lengths = []
|
||||
total = 0
|
||||
for seg in segments:
|
||||
total += len(seg)
|
||||
self.cum_lengths.append(total)
|
||||
self.total_length = total if segments else 0
|
||||
|
||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||
if not (0 <= begin_idx < self.total_length and 0 <= end_idx <= self.total_length):
|
||||
raise ValueError("begin_idx or end_idx out of bounds")
|
||||
if begin_idx >= end_idx:
|
||||
return torch.tensor([], dtype=torch.long)
|
||||
|
||||
seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx - 1)
|
||||
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx - 1)
|
||||
|
||||
result_segments = []
|
||||
|
||||
for i in range(seg_start_idx, seg_end_idx + 1):
|
||||
prev_cum = self.cum_lengths[i - 1] if i > 0 else 0
|
||||
start = max(begin_idx - prev_cum, 0)
|
||||
end = min(end_idx - prev_cum, len(self.segments[i]))
|
||||
result_segments.append(self.segments[i][start:end])
|
||||
|
||||
return torch.cat(result_segments, dim=0)
|
||||
|
||||
|
||||
class MutiSegmentFetcher:
|
||||
def __init__(self, muti_segments: MutiSeg):
|
||||
self.muti_keys = list(muti_segments.keys())
|
||||
self.muti_fetchers = {
|
||||
key: BaseSegmentFetcher(segments)
|
||||
for key, segments in muti_segments.items()
|
||||
}
|
||||
|
||||
def key_fetch(self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]) -> Union[Tensor, Seg]:
|
||||
fetch_dict = {}
|
||||
keys = [keys] if isinstance(keys, str) else keys
|
||||
|
||||
for key in keys:
|
||||
fetcher = self.muti_fetchers[key]
|
||||
fetch_tensor = fetcher.fetch_data(begin_idx, end_idx)
|
||||
fetch_dict[key] = fetch_tensor
|
||||
|
||||
return fetch_dict if len(keys) > 1 else fetch_dict[keys[0]]
|
||||
|
||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Union[Tensor, Seg]:
|
||||
return self.key_fetch(begin_idx, end_idx, self.muti_keys)
|
||||
|
||||
|
||||
class BaseDataset(Dataset, ABC):
|
||||
def __init__(self, chunk_size: int, device: str):
|
||||
super().__init__()
|
||||
self.segments: MutiSeg = {}
|
||||
self.chunk_size = chunk_size
|
||||
self.total_samples = 0
|
||||
self.device = device
|
||||
|
||||
def save(self, save_path: str):
|
||||
first_item = self.segments[keys[0]]
|
||||
segment_size = len(first_item)
|
||||
keys = list(self.segments.keys())
|
||||
|
||||
for i in range(segment_size):
|
||||
formated_segment = {key: self.segments[key][i] for key in keys}
|
||||
pkl.dump(formated_segment, open(f"{save_path}_{i}.pkl", "wb"))
|
||||
|
||||
|
||||
def load(self, load_path: Union[str, List[str]]):
|
||||
paths = [load_path] if isinstance(load_path, str) else load_path
|
||||
self.segments, self.total_samples = load_pkl_files(paths)
|
||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
||||
|
||||
@abstractmethod
|
||||
def __getitem__(self, index: int):
|
||||
raise NotImplementedError
|
||||
|
||||
def __len__(self) -> int:
|
||||
assert self.total_samples // self.chunk_size > 0
|
||||
return self.total_samples // self.chunk_size
|
||||
|
||||
|
||||
class SeqDataset(BaseDataset):
|
||||
def __init__(self, chunk_size , device='cuda'):
|
||||
super().__init__(chunk_size, device)
|
||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, "sequence")
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx = index * self.chunk_size
|
||||
end_idx = min(begin_idx + self.chunk_size, self.total_samples - 1)
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx).to(device=self.device, dtype=torch.long)
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(device=self.device, dtype=torch.long)
|
||||
|
||||
return x, y
|
||||
|
||||
|
||||
class SftDataset(BaseDataset):
|
||||
def __init__(self, chunk_size, device='cuda'):
|
||||
super().__init__(chunk_size, device)
|
||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx = index * self.chunk_size
|
||||
end_idx = min(begin_idx + self.chunk_size, self.total_samples - 1)
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx, "sequence").to(device=self.device, dtype=torch.long)
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(device=self.device, dtype=torch.long)
|
||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "mask").to(device=self.device, dtype=torch.bool)
|
||||
|
||||
return x, y, loss_mask
|
||||
|
||||
|
||||
class DpoDataset(BaseDataset):
|
||||
def __init__(self, chunk_size: int, device="cuda"):
|
||||
super().__init__(chunk_size, device)
|
||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
start_idx = index * self.chunk_size
|
||||
end_idx = min(start_idx + self.chunk_size, self.total_samples - 1)
|
||||
|
||||
chosen = self._fetch_data(start_idx, end_idx, "chosen").to(device=self.device, dtype=torch.long)
|
||||
rejected = self._fetch_data(start_idx, end_idx, "rejected").to(device=self.device, dtype=torch.long)
|
||||
chosen_mask = self._fetch_data(start_idx, end_idx, "chosen_mask").to(device=self.device, dtype=torch.bool)
|
||||
rejected_mask = self._fetch_data(start_idx, end_idx, "rejected_mask").to(device=self.device, dtype=torch.bool)
|
||||
|
||||
return chosen, rejected, chosen_mask, rejected_mask
|
||||
|
||||
|
||||
class PpoDataset(BaseDataset):
|
||||
def __init__(self, chunk_size: int, device="cuda"):
|
||||
super().__init__(chunk_size, device)
|
||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
|
||||
begin_idx = index * self.chunk_size
|
||||
end_idx = min(begin_idx + self.chunk_size, self.total_samples - 1)
|
||||
|
||||
|
||||
input_ids = self._fetch_data(begin_idx, end_idx, "input_ids").to(self.device),
|
||||
actions = self._fetch_data(begin_idx, end_idx, "actions").to(self.device),
|
||||
logprobs = self._fetch_data(begin_idx, end_idx, "logprobs").to(self.device),
|
||||
rewards = self._fetch_data(begin_idx, end_idx, "rewards").to(self.device)
|
||||
|
||||
return input_ids, actions, logprobs, rewards
|
||||
|
||||
|
||||
class DatasetLoader:
|
||||
@staticmethod
|
||||
def load(
|
||||
train_type: Literal["seq", "sft", "dpo"],
|
||||
load_path: Union[str, List[str]],
|
||||
max_len: int,
|
||||
device: str
|
||||
) -> BaseDataset:
|
||||
|
||||
dataset_router: Dict[str, Callable[[int, torch.device], BaseDataset]] = {
|
||||
"seq": lambda m_len, device: SeqDataset(m_len, device=device),
|
||||
"sft": lambda m_len, device: SftDataset(m_len, device=device),
|
||||
"dpo": lambda m_len, device: DpoDataset(m_len, device=device),
|
||||
}
|
||||
dataset = dataset_router[train_type](max_len, device)
|
||||
dataset.load(load_path)
|
||||
|
||||
return dataset
|
||||
@@ -0,0 +1,55 @@
|
||||
|
||||
import torch
|
||||
from abc import abstractmethod
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
|
||||
class MaskBuilder:
|
||||
def __init__(
|
||||
self,
|
||||
bos_token_id: int,
|
||||
eos_token_id: int,
|
||||
user_token_id: int,
|
||||
system_token_id: int,
|
||||
|
||||
):
|
||||
self.bos_token_id = bos_token_id
|
||||
self.eos_token_id = eos_token_id
|
||||
self.user_token_id = user_token_id
|
||||
self.system_token_id = system_token_id
|
||||
|
||||
@abstractmethod
|
||||
def build(input_ids: Tensor) -> Tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
|
||||
class LossMaskBuilder(MaskBuilder):
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def build(self, input_ids: Tensor) -> Tensor:
|
||||
token_markers = torch.zeros_like(input_ids, dtype=torch.int8)
|
||||
|
||||
is_user_token = input_ids.eq(self.user_token_id)
|
||||
is_system_token = input_ids.eq(self.system_token_id)
|
||||
|
||||
token_markers[is_user_token] = 1
|
||||
token_markers[is_system_token] = -1
|
||||
|
||||
cumulative_markers = torch.cumsum(token_markers, dim=-1)
|
||||
min_cumulative = cumulative_markers.min(dim=-1, keepdim=True).values
|
||||
loss_mask = cumulative_markers - min_cumulative
|
||||
|
||||
return loss_mask
|
||||
|
||||
|
||||
|
||||
|
||||
class AttentionMaskBuilder:
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def build(input_ids: Tensor):
|
||||
bsz = input_ids.size(0)
|
||||
@@ -0,0 +1,388 @@
|
||||
import copy
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from torch import Tensor
|
||||
from torch.optim import Optimizer
|
||||
from torch.utils.data import Dataset
|
||||
from typing import Any, Literal, Tuple, Callable, Dict
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import asdict, dataclass, field
|
||||
|
||||
|
||||
def get_logprobs(model:nn.Module, input_ids: Tensor, mask: Tensor, pad_token_id):
|
||||
input_mask = input_ids.ne(pad_token_id)
|
||||
logits = model(input_ids, input_mask)["logits"]
|
||||
log_probs = torch.log_softmax(logits, dim=-1)
|
||||
|
||||
shifted_log_probs = log_probs[:, :-1, :]
|
||||
shifted_input_ids = input_ids[:, 1:]
|
||||
shifted_response_mask = mask[:, 1:]
|
||||
|
||||
token_logprobs = torch.gather(
|
||||
shifted_log_probs,
|
||||
dim=-1,
|
||||
index=shifted_input_ids.unsqueeze(-1)
|
||||
).squeeze(-1)
|
||||
|
||||
prompt_mask = input_mask[:, 1:]
|
||||
valid_mask = (prompt_mask & shifted_response_mask).float()
|
||||
|
||||
return (token_logprobs * valid_mask).sum(dim=-1)
|
||||
|
||||
|
||||
class MaskBuilder:
|
||||
def __init__(
|
||||
self,
|
||||
bos_token_id: int,
|
||||
eos_token_id: int,
|
||||
user_token_id: int,
|
||||
system_token_id: int,
|
||||
|
||||
):
|
||||
self.bos_token_id = bos_token_id
|
||||
self.eos_token_id = eos_token_id
|
||||
self.user_token_id = user_token_id
|
||||
self.system_token_id = system_token_id
|
||||
|
||||
@abstractmethod
|
||||
def build(input_ids: Tensor) -> Tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
|
||||
class LossMaskBuilder(MaskBuilder):
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def build(self, input_ids: Tensor) -> Tensor:
|
||||
token_markers = torch.zeros_like(input_ids, dtype=torch.int8)
|
||||
|
||||
is_user_token = input_ids.eq(self.user_token_id)
|
||||
is_system_token = input_ids.eq(self.system_token_id)
|
||||
|
||||
token_markers[is_user_token] = 1
|
||||
token_markers[is_system_token] = -1
|
||||
|
||||
cumulative_markers = torch.cumsum(token_markers, dim=-1)
|
||||
min_cumulative = cumulative_markers.min(dim=-1, keepdim=True).values
|
||||
loss_mask = cumulative_markers - min_cumulative
|
||||
|
||||
return loss_mask.to(dtype=torch.bool)
|
||||
|
||||
|
||||
|
||||
|
||||
class AttentionMaskBuilder(MaskBuilder):
|
||||
def __init__(self, multi_turn=False, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.multi_turn = multi_turn
|
||||
|
||||
def build(self, input_ids: Tensor):
|
||||
bsz = input_ids.size(0)
|
||||
|
||||
|
||||
def _build_batch(self, input_ids: Tensor):
|
||||
is_user_token = input_ids.eq(self.user_token_id)
|
||||
|
||||
token_markers = torch.zeros_like(input_ids, dtype=torch.int8)
|
||||
token_markers[is_user_token] = 1
|
||||
cumulative_markers = torch.cumsum(token_markers, dim=-1)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class BaseStrategy(ABC):
|
||||
def __init__(self, model: nn.Module):
|
||||
self.model = model
|
||||
|
||||
@abstractmethod
|
||||
def compute_loss(self, batch: Tuple[Tensor, ...]) -> Tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
def __call__(self, batch: Tuple[Tensor, ...]) -> Tensor:
|
||||
return self.compute_loss(batch)
|
||||
|
||||
|
||||
class SeqStrategy(BaseStrategy):
|
||||
def __init__(self, model):
|
||||
super().__init__(model)
|
||||
|
||||
def compute_loss(self, batch: Tuple[Tensor, ...]) -> Tensor:
|
||||
x, y = batch
|
||||
B, L = x.size()
|
||||
logits: Tensor = self.model(x)["logits"]
|
||||
|
||||
loss = F.cross_entropy(
|
||||
logits.view(B * L, -1), y.flatten()
|
||||
)
|
||||
return loss
|
||||
|
||||
|
||||
class SftStrategy(BaseStrategy):
|
||||
def __init__(self, model):
|
||||
super().__init__(model)
|
||||
|
||||
def compute_loss(self, batch: Tuple[Tensor, ...]) -> Tensor:
|
||||
x, y, loss_mask = batch
|
||||
B, L = x.size()
|
||||
ignore_idx = -1
|
||||
|
||||
logits: Tensor = self.model(x)["logits"]
|
||||
masked_y = y.masked_fill(loss_mask == 0, ignore_idx)
|
||||
|
||||
loss = F.cross_entropy(
|
||||
logits.view(B * L, -1),
|
||||
masked_y.flatten(),
|
||||
ignore_index=ignore_idx
|
||||
)
|
||||
|
||||
return loss
|
||||
|
||||
class DpoStrategy(BaseStrategy):
|
||||
def __init__(self, model, pad_token_id, beta):
|
||||
super().__init__(model)
|
||||
ref_model = copy.deepcopy(self.model)
|
||||
ref_model.requires_grad_(False)
|
||||
ref_model.eval()
|
||||
|
||||
self.ref_model = ref_model
|
||||
self.pad_token_id = pad_token_id
|
||||
self.beta = beta
|
||||
|
||||
def compute_loss(self, batch: Tuple[Tensor, ...]) -> Tensor:
|
||||
good_ids, bad_ids, good_mask, bad_mask = batch
|
||||
|
||||
log_pi_good = get_logprobs(self.model, good_ids, good_mask, self.pad_token_id)
|
||||
log_pi_bad = get_logprobs(self.model, bad_ids, bad_mask, self.pad_token_id)
|
||||
|
||||
with torch.no_grad():
|
||||
log_ref_good = get_logprobs(self.ref_model, good_ids, good_mask, self.pad_token_id)
|
||||
log_ref_bad = get_logprobs(self.ref_model, bad_ids, bad_mask, self.pad_token_id)
|
||||
|
||||
pi_log_ratio = log_pi_good - log_pi_bad
|
||||
ref_log_ratio = log_ref_good - log_ref_bad
|
||||
|
||||
ratio_diff = pi_log_ratio - ref_log_ratio
|
||||
|
||||
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
|
||||
return dpo_loss
|
||||
|
||||
|
||||
class PpoStrategy(BaseStrategy):
|
||||
def __init__(self, model, pad_token_id, epsilon):
|
||||
super().__init__(model)
|
||||
ref_model = copy.deepcopy(self.model)
|
||||
ref_model.requires_grad_(False)
|
||||
ref_model.eval()
|
||||
|
||||
self.ref_model = ref_model
|
||||
self.pad_token_id = pad_token_id
|
||||
self.epsilon = epsilon
|
||||
|
||||
def ppo_clip_loss_masked(
|
||||
self,
|
||||
log_probs: Tensor,
|
||||
old_log_probs: Tensor,
|
||||
advantages: Tensor,
|
||||
values: Tensor,
|
||||
returns: Tensor,
|
||||
mask: Tensor,
|
||||
clip_eps: float=0.2,
|
||||
):
|
||||
ratio = torch.exp(log_probs - old_log_probs)
|
||||
surr1 = ratio * advantages
|
||||
surr2 = torch.clamp(ratio, 1 - clip_eps, 1 + clip_eps) * advantages
|
||||
policy_loss = -torch.min(surr1, surr2).masked_select(mask).mean()
|
||||
|
||||
value_loss = F.mse_loss(values.masked_select(mask),
|
||||
returns.masked_select(mask))
|
||||
|
||||
entropy = -(log_probs.exp() * log_probs).masked_select(mask).mean()
|
||||
entropy_loss = -entropy
|
||||
return policy_loss, value_loss, entropy_loss
|
||||
|
||||
|
||||
|
||||
class StrategyFactory:
|
||||
|
||||
def load(model, train_type, pad_token_id, dpo_beta):
|
||||
train_strategy: Dict[str, Callable[[], BaseStrategy]] = {
|
||||
"seq": lambda: SeqStrategy(model),
|
||||
"sft": lambda: SftStrategy(model),
|
||||
"dpo": lambda: DpoStrategy(model, pad_token_id, dpo_beta)
|
||||
}
|
||||
strategy = train_strategy[train_type]()
|
||||
return strategy
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrainConfig:
|
||||
train_type: str = field(
|
||||
default_factory=["seq", "sft", "dpo"],
|
||||
metadata={"help": "Type of training."}
|
||||
)
|
||||
dataset: Dataset = field(
|
||||
default=None,
|
||||
metadata={"help": "Dataset for training."}
|
||||
)
|
||||
optimizer: Optimizer = field(
|
||||
default=None,
|
||||
metadata={"help": "Optimizer for training."}
|
||||
)
|
||||
ckpt_dir: str = field(
|
||||
default="./checkpoint",
|
||||
metadata={"help": "Checkpoint directory."}
|
||||
)
|
||||
n_epoch: int = field(
|
||||
default=1,
|
||||
metadata={"help": "Number of epochs for training."}
|
||||
)
|
||||
batch_size: int = field(
|
||||
default=4,
|
||||
metadata={"help": "Batch size for training."}
|
||||
)
|
||||
n_iter_ckpt: int = field(
|
||||
default=5000,
|
||||
metadata={"help": "Number of iterations between checkpoints."}
|
||||
)
|
||||
n_iter_step: int = field(
|
||||
default=1,
|
||||
metadata={"help": "Number of iterations between steps."}
|
||||
)
|
||||
max_grad_norm: float = field(
|
||||
default=1.0,
|
||||
metadata={"help": "Maximum gradient norm."}
|
||||
)
|
||||
random_seed: int = field(
|
||||
default=3407,
|
||||
metadata={"help": "Random seed."}
|
||||
)
|
||||
dpo_beta: float = field(
|
||||
default=0.1,
|
||||
metadata={"help": "DPO beta."}
|
||||
)
|
||||
|
||||
def get_kwargs(self)-> Dict[str, Any]:
|
||||
config_dict = asdict(self)
|
||||
return {k: v for k, v in config_dict.items() if v is not None}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScheduleConfig:
|
||||
schedule_type: str = field(
|
||||
default_factory=["cosine", "sgdr"],
|
||||
metadata={"help": "Type of learning rate schedule."}
|
||||
)
|
||||
warning_step: int = field(
|
||||
default=1000,
|
||||
metadata= {"help": "Warning up step."}
|
||||
)
|
||||
@abstractmethod
|
||||
def get_kwargs(self)-> Dict[str, Any]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclass
|
||||
class CosineScheduleConfig(ScheduleConfig):
|
||||
total_iters: int = field(
|
||||
default=None,
|
||||
metadata={"help": "Total iterations for cosine schedule."}
|
||||
)
|
||||
min_rate: float = field(
|
||||
default=0.05,
|
||||
metadata={"help": "Minimum rate for cosine schedule."}
|
||||
)
|
||||
schedule_type: Literal["cosine"] = "cosine"
|
||||
|
||||
def get_kwargs(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"schedule_type": self.schedule_type,
|
||||
"warning_step": self.warning_step,
|
||||
"lr_decay_iters": self.total_iters - self.warning_step,
|
||||
"min_rate": self.min_rate
|
||||
}
|
||||
|
||||
@dataclass
|
||||
class SgdrScheduleConfig(ScheduleConfig):
|
||||
cycle_length: int = field(
|
||||
default=1000,
|
||||
metadata={"help": "Cycle length for sgdr schedule."}
|
||||
)
|
||||
min_rate: float = field(
|
||||
default=0.05,
|
||||
metadata={"help": "Minimum rate for sgdr schedule."}
|
||||
)
|
||||
T_mult: int = field(
|
||||
default=2,
|
||||
metadata={"help": "T_mult for sgdr schedule."}
|
||||
)
|
||||
schedule_type: Literal["sgdr"] = "sgdr"
|
||||
|
||||
def get_kwargs(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"schedule_type": self.schedule_type,
|
||||
"warning_step": self.warning_step,
|
||||
"cycle_length": self.cycle_length,
|
||||
"min_rate": self.min_rate,
|
||||
"T_mult": self.T_mult
|
||||
}
|
||||
|
||||
|
||||
class SchedulerFactory:
|
||||
|
||||
@staticmethod
|
||||
def get_sgdr_schedule(
|
||||
warning_step: int,
|
||||
cycle_length: int,
|
||||
min_rate: float = 0.1,
|
||||
T_mult: int = 2
|
||||
) -> Callable[[int], float]:
|
||||
|
||||
def sgdr_schedule(now_iter: int) -> float:
|
||||
if now_iter < warning_step:
|
||||
return max(min_rate, now_iter / warning_step)
|
||||
|
||||
adjusted_iter = now_iter - warning_step
|
||||
total_cycles, current_cycle = 0, 0
|
||||
while adjusted_iter >= cycle_length * (T_mult ** total_cycles):
|
||||
current_cycle += 1
|
||||
total_cycles += 1
|
||||
|
||||
cycle_start = sum(cycle_length * (T_mult ** i) for i in range(current_cycle))
|
||||
cycle_pos = adjusted_iter - cycle_start
|
||||
|
||||
cycle_length_current = cycle_length * (T_mult ** current_cycle)
|
||||
return max(min_rate, 0.5 * (1 + math.cos(math.pi * cycle_pos / cycle_length_current)))
|
||||
|
||||
return sgdr_schedule
|
||||
|
||||
@staticmethod
|
||||
def get_cosine_warmup_schedule(
|
||||
warning_step: int,
|
||||
lr_decay_iters: int,
|
||||
min_rate: float = 0.1
|
||||
) -> Callable[[int], float]:
|
||||
|
||||
def cosine_warmup_schedule(now_iter: int) -> float:
|
||||
if now_iter <= warning_step:
|
||||
return max(min_rate, now_iter / warning_step)
|
||||
else:
|
||||
rate = (now_iter - warning_step) / (lr_decay_iters - warning_step)
|
||||
return max(min_rate, 0.5 * (1.0 + math.cos(math.pi * rate)))
|
||||
|
||||
return cosine_warmup_schedule
|
||||
|
||||
@staticmethod
|
||||
def load_schedule_fn(**kwargs):
|
||||
strategy = kwargs.pop("schedule_type")
|
||||
if strategy == "cosine":
|
||||
return SchedulerFactory.get_cosine_warmup_schedule(**kwargs)
|
||||
elif strategy == "sgdr":
|
||||
return SchedulerFactory.get_sgdr_schedule(**kwargs)
|
||||
else:
|
||||
raise ValueError(f"Invalid schedule type: {strategy}")
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
import os
|
||||
import torch
|
||||
import logging
|
||||
|
||||
from typing import Tuple
|
||||
from torch.nn.utils import clip_grad_norm_
|
||||
from torch.optim.lr_scheduler import LambdaLR
|
||||
from torch.utils.data import DataLoader, RandomSampler
|
||||
from tqdm import tqdm
|
||||
|
||||
from khaosz.core import ModelParameter, Checkpoint
|
||||
from khaosz.trainer.strategy import SchedulerFactory, StrategyFactory, TrainConfig, ScheduleConfig
|
||||
|
||||
|
||||
class Trainer:
|
||||
def __init__(
|
||||
self,
|
||||
parameter: ModelParameter,
|
||||
log_path: str="./train_log.log"
|
||||
):
|
||||
logger = logging.getLogger()
|
||||
logger.setLevel(level = logging.INFO)
|
||||
handler = logging.FileHandler(log_path)
|
||||
handler.setLevel(logging.INFO)
|
||||
handler.setFormatter(logging.Formatter('%(asctime)s: %(message)s'))
|
||||
logger.addHandler(handler)
|
||||
logger.info("initializing trainer ...")
|
||||
|
||||
self.logger = logger
|
||||
self.model = parameter.model
|
||||
self.tokenizer = parameter.tokenizer
|
||||
self.config = parameter.config
|
||||
|
||||
def save_checkpoint(
|
||||
self,
|
||||
loss_list: list,
|
||||
ckpt_dir: str,
|
||||
current_iter: int,
|
||||
last_ckpt_iter: int
|
||||
):
|
||||
save_path = os.path.join(ckpt_dir, f"iter_{current_iter}")
|
||||
Checkpoint(
|
||||
self.model,
|
||||
self.tokenizer,
|
||||
self.config,
|
||||
loss_list,
|
||||
current_iter
|
||||
).save(save_path)
|
||||
|
||||
diff_iter = current_iter - last_ckpt_iter
|
||||
avg_loss = sum(loss_list[last_ckpt_iter:current_iter]) / diff_iter
|
||||
self.logger.info(f"iter: {current_iter} loss: {avg_loss}")
|
||||
|
||||
return current_iter
|
||||
|
||||
def load_checkpoint(self, train_checkpoint: Checkpoint) -> Tuple[list, int]:
|
||||
self.model = train_checkpoint.model
|
||||
self.tokenizer = train_checkpoint.tokenizer
|
||||
self.config = train_checkpoint.config
|
||||
loss_list = train_checkpoint.loss_list
|
||||
last_ckpt_iter = train_checkpoint.current_iter
|
||||
|
||||
return loss_list, last_ckpt_iter
|
||||
|
||||
def train(
|
||||
self,
|
||||
train_config: TrainConfig,
|
||||
schedule_config: ScheduleConfig,
|
||||
train_checkpoint: Checkpoint = None
|
||||
):
|
||||
assert schedule_config.schedule_type in ["cosine", "sgdr"]
|
||||
assert train_config.train_type in ["seq", "sft", "dpo"]
|
||||
|
||||
if train_checkpoint:
|
||||
loss_list, last_ckpt_iter = self.load_checkpoint(train_checkpoint)
|
||||
current_iter = train_checkpoint.current_iter + 1
|
||||
self.logger.info(f"Resuming training from checkpoint: iter {current_iter}")
|
||||
else:
|
||||
current_iter = 0
|
||||
last_ckpt_iter = 0
|
||||
loss_list = []
|
||||
|
||||
lambda_scheduler_fn = SchedulerFactory.load_schedule_fn(
|
||||
**schedule_config.get_kwargs()
|
||||
)
|
||||
|
||||
strategy = StrategyFactory.load(
|
||||
self.model,
|
||||
train_config.train_type,
|
||||
self.tokenizer.pad_id,
|
||||
train_config.dpo_beta
|
||||
)
|
||||
|
||||
scheduler = LambdaLR(
|
||||
train_config.optimizer,
|
||||
lambda_scheduler_fn,
|
||||
last_epoch=current_iter - 1 if train_checkpoint else -1
|
||||
)
|
||||
|
||||
seed = train_config.random_seed
|
||||
generator = torch.Generator().manual_seed(seed)
|
||||
sampler = RandomSampler(train_config.dataset, generator=generator)
|
||||
remaining_epochs = train_config.n_epoch - current_iter // (len(train_config.dataset) // train_config.batch_size)
|
||||
|
||||
self.logger.info(f"Starting {train_config.train_type.upper()} training for {train_config.n_epoch} epochs")
|
||||
self.logger.info(f"Checkpoint interval: {train_config.n_iter_ckpt} iterations")
|
||||
|
||||
for epoch in range(remaining_epochs):
|
||||
self.model.train()
|
||||
dataloader = DataLoader(
|
||||
train_config.dataset,
|
||||
batch_size=train_config.batch_size,
|
||||
sampler=sampler
|
||||
)
|
||||
progress_bar = tqdm(
|
||||
dataloader,
|
||||
desc=f"Epoch {epoch+1}/{train_config.n_epoch}",
|
||||
dynamic_ncols=True
|
||||
)
|
||||
for batch in progress_bar:
|
||||
#forward
|
||||
loss = strategy(batch)
|
||||
loss_list.append(loss.item())
|
||||
#backward
|
||||
loss.backward()
|
||||
#step
|
||||
if current_iter % train_config.n_iter_step == 0:
|
||||
clip_grad_norm_(
|
||||
self.model.parameters(),
|
||||
train_config.max_grad_norm
|
||||
)
|
||||
train_config.optimizer.step()
|
||||
train_config.optimizer.zero_grad()
|
||||
|
||||
current_iter += 1
|
||||
scheduler.step()
|
||||
progress_bar.set_postfix({
|
||||
"loss": f"{loss.item():.4f}",
|
||||
"lr": f"{train_config.optimizer.param_groups[0]['lr']:.2e}"
|
||||
})
|
||||
#save checkpotint
|
||||
if current_iter - last_ckpt_iter >= train_config.n_iter_ckpt:
|
||||
last_ckpt_iter = self.save_checkpoint(
|
||||
loss_list,
|
||||
train_config.ckpt_dir,
|
||||
current_iter,
|
||||
last_ckpt_iter
|
||||
)
|
||||
|
||||
if current_iter != last_ckpt_iter:
|
||||
last_ckpt_iter = self.save_checkpoint(
|
||||
loss_list,
|
||||
train_config.ckpt_dir,
|
||||
current_iter,
|
||||
last_ckpt_iter
|
||||
)
|
||||
|
||||
self.logger.info("Training completed")
|
||||
|
||||
return Checkpoint(
|
||||
self.model,
|
||||
self.tokenizer,
|
||||
self.config,
|
||||
loss_list,
|
||||
current_iter,
|
||||
train_config.optimizer
|
||||
)
|
||||
Reference in New Issue
Block a user