Initial commit

This commit is contained in:
2025-09-27 12:02:22 +08:00
commit a4443765ee
33 changed files with 3896 additions and 0 deletions
+11
View File
@@ -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",
]
+210
View File
@@ -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
+55
View File
@@ -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)
+388
View File
@@ -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}")
+167
View File
@@ -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
)