58 Commits
Author SHA1 Message Date
ViperEkura f31bf5a959 test(transformer): 更新 tie_weight 相关测试逻辑 2025-11-09 17:23:33 +08:00
ViperEkura 7a21f5d72e build(setup): 更新版本号并调整 Python 版本要求 2025-11-09 16:40:20 +08:00
ViperEkura 0b45e8666e fix(scripts): 修复stream_chat.py中的拼写错误 2025-11-09 16:30:24 +08:00
ViperEkura 6f3386f02c fix(transformer): 优化state_dict 处理逻辑, 优化attention_mask的处理方式 2025-11-09 16:25:17 +08:00
ViperEkura d25202a329 feat(model): 实现旋转位置编码缓存动态扩展 2025-11-09 14:35:29 +08:00
ViperEkura 254ec934be feat(transformer): 简化权重绑定逻辑并增加测试单元 2025-11-07 15:14:54 +08:00
ViperEkura 7e5ecf3b7d refactor(config): 重命名 TransformerConfig 为 ModelConfig 2025-11-07 07:31:12 +08:00
ViperEkura 66a551217e refactor(generator): 优化生成逻辑 2025-11-07 07:24:00 +08:00
ViperEkura bdc3f4dc63 feat(module): 重构旋转位置编码实现以提升性能和可读性 2025-11-06 17:52:47 +08:00
ViperEkura 805773c7fe docs(transformer): 更新process_attention_mask函数文档 2025-11-05 23:41:11 +08:00
ViperEkura 7ccc4ab9ac fix(model): 修复加载状态字典时的权重共享问题 2025-11-05 23:38:45 +08:00
ViperEkura 69d9374f51 feat(model): 添加 tie_weight 配置选项并优化模型模块实现 2025-11-05 23:26:57 +08:00
ViperEkura b260f5581d fix(benchmark): 优化 KV 缓存初始化并更正基准测试类型标识 2025-11-05 15:44:29 +08:00
ViperEkura 0a754e3341 feat(scripts): 调整文本生成参数以提升多样性 2025-11-05 13:56:58 +08:00
ViperEkura 144b9598ad feat(model): 添加 Linear 和 Embedding 模块的自定义参数初始化支持 2025-10-31 22:43:12 +08:00
ViperEkura 877669b799 feat(inference): 添加generate_loop方法并优化KVCacheManager初始化 2025-10-31 21:15:15 +08:00
ViperEkura cdb47a62dc test: 统一重构数据集和调度器测试模块 2025-10-31 20:24:01 +08:00
ViperEkura e86328b753 fix(tokenizer): 修复stop_ids属性返回错误的token ID列表 2025-10-31 19:19:38 +08:00
ViperEkura 5d3799b715 refactor(data): 修改变量命名方式 2025-10-30 16:32:25 +08:00
ViperEkura 6a3135f401 fix(data_util): 修复数据集索引计算逻辑并提取通用方法 2025-10-29 20:58:33 +08:00
ViperEkura 12850d403c fix(config): 修改Checkpoint类中tokenizer和config字段的默认值初始化方式 2025-10-29 13:24:20 +08:00
ViperEkura bad6243b53 fix(train): 更新训练函数参数传递方式 2025-10-29 13:23:53 +08:00
ViperEkura f2448a5147 feat(benchmark): 优化KV缓存初始化逻辑 2025-10-29 12:41:32 +08:00
ViperEkura 46b2a0f86f feat(train): 添加 max_len 和 step_size 参数支持 2025-10-29 12:32:17 +08:00
ViperEkura d94fc5a87a feat(data, inference): 使用chatML格式 2025-10-29 12:02:43 +08:00
ViperEkura 38b2725cd1 feat(KVCacheManager): 优化KV缓存结构为元组形式以提升性能 2025-10-29 12:01:28 +08:00
ViperEkura bc5ef72001 fix(config): 修正 SGDRScheduleConfig 类名拼写错误 2025-10-20 18:21:46 +08:00
ViperEkura e051005334 test(test_module): 更新测试用例以使用新的generate_iterator接口 2025-10-20 13:52:31 +08:00
ViperEkura 0db046f8d9 feat(khaosz/trainer): 更新梯度裁剪回调 2025-10-20 13:30:26 +08:00
ViperEkura 05b012820b refactor(khaosz): 重构模块导出结构并重命名主模块文件 2025-10-20 13:07:02 +08:00
ViperEkura e72e244df6 feat(inference): 实现采样策略并优化生成器逻辑 2025-10-20 13:00:41 +08:00
ViperEkura 98efca7b9d feat(trainer): 添加训练起始轮次和批次配置支持 2025-10-19 21:47:10 +08:00
ViperEkura 613edd7a14 test(early_stopping, train_strategy): 更新测试配置以提高稳定性 2025-10-18 22:07:11 +08:00
ViperEkura 622982364b fix(trainer): 修复检查点加载逻辑 2025-10-18 21:45:23 +08:00
ViperEkura b67bc9865d refactor(trainer): 重构学习率调度器实现并分离配置与工厂逻辑 2025-10-18 16:42:37 +08:00
ViperEkura c51b203fde refactor(khaosz): 重构项目结构 2025-10-18 13:56:59 +08:00
ViperEkura 8434c19923 fix(khaosz/trainer): 修复数据获取中的索引范围错误和参数传递问题 2025-10-09 19:53:52 +08:00
ViperEkura 68a15005cb feat(train.py): 支持从检查点恢复训练并优化数据加载配置 2025-10-07 22:02:50 +08:00
ViperEkura efbe3de9d3 fix(khaosz/trainer/data_util): 修复数据集索引范围错误 2025-10-07 20:04:45 +08:00
ViperEkura 12793bc2d3 feat(khaosz/trainer): 新增梯度统计工具函数并重构训练回调机制 2025-10-07 13:03:32 +08:00
ViperEkura 0764cb8296 fix(khaosz/trainer/train_callback): 修复基类函数命名错误 2025-10-07 11:43:51 +08:00
ViperEkura 57cd7b921e feat(khaosz/trainer): 改进训练循环中的损失归一化处理 2025-10-06 20:17:47 +08:00
ViperEkura c1bf22b6ec refactor(khaosz/trainer): 使用 TrainContext 替代 kwargs 传递训练上下文 2025-10-06 20:12:08 +08:00
ViperEkura f9b6331ad7 refactor(khaosz/core/parameter): 修改参数名称 2025-10-06 20:11:46 +08:00
ViperEkura 183f481692 build(khaosz): 更新版本号至1.3.0 2025-10-06 17:12:12 +08:00
ViperEkura ec0c054d26 test(early_stopping): 移除未使用的torch.utils.data导入 2025-10-06 17:10:10 +08:00
ViperEkura 4ffa7454f2 feat(strategy): 支持模型输入可调用对象并优化损失计算 2025-10-06 17:08:56 +08:00
ViperEkura 8c9e973179 fix(train.py): 修复数据集加载时的参数传递问题 2025-10-06 16:44:02 +08:00
ViperEkura fc98d9b7e6 refactor(khaosz/trainer): 移除未使用的导入模块 2025-10-04 21:45:53 +08:00
ViperEkura 9d5aa952e0 feat(tests): 重构测试环境, 便于pickle 序列化 2025-10-04 21:31:39 +08:00
ViperEkura 2ccd7bd583 refactor(khaosz/trainer): 重构训练器模块结构以提升可维护性 2025-10-04 21:31:15 +08:00
ViperEkura e7d29ca2d5 feat(tests): 改进测试环境配置与设备管理 2025-10-04 12:12:42 +08:00
ViperEkura 465a1a9373 refactor(khaosz/tainer): 修改设备参数传递发生阶段 2025-10-04 12:12:21 +08:00
ViperEkura 240ee00221 feat(khaosz/trainer): 引入 TrainContext 和 TrainContextBuilder 优化训练上下文管理 2025-10-03 22:42:11 +08:00
ViperEkura 6e1a497c04 test(sampler): 删除冗余的训练恢复测试用例 2025-10-03 22:18:31 +08:00
ViperEkura 85aeec9e55 test(conftest): 添加matplotlib后端设置以避免GUI问题 2025-10-03 22:11:54 +08:00
ViperEkura 9a452dd34e fix(khaosz/trainer/data_util.py): 修复 RandomSampler 中迭代计数器位置错误 2025-10-03 22:08:28 +08:00
ViperEkura 28b01220b6 test(trainer): 拆分测试文件 2025-10-03 22:08:11 +08:00
43 changed files with 2750 additions and 2245 deletions
+14 -31
View File
@@ -1,7 +1,7 @@
import torch import torch
from typing import Dict, Any from typing import Dict, Any
from dataclasses import dataclass from dataclasses import dataclass
from khaosz.core.transformer import TransformerConfig, Transformer from khaosz.model.transformer import ModelConfig, Transformer
@dataclass @dataclass
@@ -15,7 +15,7 @@ class BenchmarkResult:
class GenerationBenchmark: class GenerationBenchmark:
def __init__( def __init__(
self, self,
config: TransformerConfig, config: ModelConfig,
device: str = "cuda", device: str = "cuda",
dtype: torch.dtype = torch.float16 dtype: torch.dtype = torch.float16
): ):
@@ -25,21 +25,13 @@ class GenerationBenchmark:
self.model = Transformer(config).to(device=device, dtype=dtype) self.model = Transformer(config).to(device=device, dtype=dtype)
self.model.eval() self.model.eval()
def _initialize_kv_cache(self, batch_size: int, max_len: int) -> list: def _initialize_kv_cache(self, batch_size: int) -> list:
"""初始化KV缓存""" """初始化KV缓存"""
kv_cache = [] config = self.config
head_dim = self.config.n_dim // self.config.n_head shape = (batch_size, config.n_layer, config.m_len, config.n_kvhead, config.n_dim // config.n_head)
for _ in range(self.config.n_layer): k_cache = torch.zeros(shape, device=self.device, dtype=self.dtype)
k_cache = torch.zeros( v_cache = torch.zeros(shape, device=self.device, dtype=self.dtype)
(batch_size, max_len, self.config.n_kvhead, head_dim), return (k_cache, v_cache)
device=self.device, dtype=self.dtype
)
v_cache = torch.zeros(
(batch_size, max_len, self.config.n_kvhead, head_dim),
device=self.device, dtype=self.dtype
)
kv_cache.append((k_cache, v_cache))
return kv_cache
def _prepare_inputs(self, batch_size: int, prompt_length: int, total_length: int): def _prepare_inputs(self, batch_size: int, prompt_length: int, total_length: int):
prompt_ids = torch.randint( prompt_ids = torch.randint(
@@ -121,7 +113,7 @@ class GenerationBenchmark:
for trial in range(num_trials): for trial in range(num_trials):
prompt_ids, gen_ids = self._prepare_inputs(batch_size, prompt_length, prompt_length + gen_length) prompt_ids, gen_ids = self._prepare_inputs(batch_size, prompt_length, prompt_length + gen_length)
kv_cache = self._initialize_kv_cache(batch_size, self.config.m_len) kv_cache = self._initialize_kv_cache(batch_size)
_ = self.model(prompt_ids, persistent_key_values=kv_cache, start_pos=0) _ = self.model(prompt_ids, persistent_key_values=kv_cache, start_pos=0)
torch.cuda.synchronize() torch.cuda.synchronize()
@@ -152,7 +144,7 @@ class GenerationBenchmark:
total_time=total_time, total_time=total_time,
tokens_per_second=total_tokens / total_time, tokens_per_second=total_tokens / total_time,
metadata={ metadata={
"benchmark_type": "generation", "benchmark_type": "decoding",
"batch_size": batch_size, "batch_size": batch_size,
"prompt_length": prompt_length, "prompt_length": prompt_length,
"gen_length": gen_length, "gen_length": gen_length,
@@ -173,7 +165,7 @@ def print_benchmark_result(result: BenchmarkResult):
if benchmark_type == "prefill": if benchmark_type == "prefill":
print(f"Batch Size: {result.metadata['batch_size']} | Prompt Length: {result.metadata['prompt_length']}") print(f"Batch Size: {result.metadata['batch_size']} | Prompt Length: {result.metadata['prompt_length']}")
elif benchmark_type == "generation": elif benchmark_type == "decoding":
print(f"Batch Size: {result.metadata['batch_size']} | Gen Length: {result.metadata['gen_length']}") print(f"Batch Size: {result.metadata['batch_size']} | Gen Length: {result.metadata['gen_length']}")
print(f"Device: {result.metadata['device']} | Dtype: {result.metadata['dtype']}") print(f"Device: {result.metadata['device']} | Dtype: {result.metadata['dtype']}")
@@ -181,7 +173,7 @@ def print_benchmark_result(result: BenchmarkResult):
if __name__ == "__main__": if __name__ == "__main__":
config = TransformerConfig( config = ModelConfig(
vocab_size=10000, vocab_size=10000,
n_dim=1536, n_dim=1536,
n_head=24, n_head=24,
@@ -198,18 +190,9 @@ if __name__ == "__main__":
print("Running Transformer Generation Benchmark") print("Running Transformer Generation Benchmark")
print("=" * 80) print("=" * 80)
prefill_result = benchmark.run_prefill_benchmark( prefill_result = benchmark.run_prefill_benchmark(batch_size=4, prompt_length=512, num_trials=5)
batch_size=4,
prompt_length=512,
num_trials=5
)
print_benchmark_result(prefill_result) print_benchmark_result(prefill_result)
gen_result = benchmark.run_decoding_benchmark( gen_result = benchmark.run_decoding_benchmark(batch_size=4, prompt_length=512, gen_length=128, num_trials=5)
batch_size=4,
prompt_length=512,
gen_length=128,
num_trials=5
)
print_benchmark_result(gen_result) print_benchmark_result(gen_result)
+27 -22
View File
@@ -1,16 +1,23 @@
__version__ = "1.2.2" __version__ = "1.3.1"
__author__ = "ViperEkura" __author__ = "ViperEkura"
from khaosz.model import Khaosz from khaosz.api import Khaosz
from khaosz.core.transformer import Transformer, TransformerConfig from khaosz.config import (
ModelConfig,
ParameterLoader,
TrainConfig,
)
from khaosz.model.transformer import Transformer
from khaosz.utils.retriever import Retriever from khaosz.utils.retriever import Retriever
from khaosz.utils.splitter import ( from khaosz.utils.splitter import (
SemanticTextSplitter, SemanticTextSplitter,
PriorityTextSplitter PriorityTextSplitter
) )
from khaosz.core.tokenizer import BpeTokenizer from khaosz.data import (
from khaosz.core.parameter import ParameterLoader DatasetLoader,
from khaosz.core.generator import ( BpeTokenizer
)
from khaosz.inference.generator import (
TextGenerator, TextGenerator,
ChatGenerator, ChatGenerator,
StreamGenerator, StreamGenerator,
@@ -18,23 +25,29 @@ from khaosz.core.generator import (
RetrievalGenerator, RetrievalGenerator,
EmbeddingEncoder EmbeddingEncoder
) )
from khaosz.trainer import ( from khaosz.trainer import (
Trainer, Trainer,
DatasetLoader,
TrainConfig,
StrategyFactory, StrategyFactory,
SchedulerFactory SchedulerFactory
) )
__all__ = [ __all__ = [
# model
"Khaosz", "Khaosz",
# module
"Transformer", "Transformer",
"TransformerConfig",
"BpeTokenizer", "Retriever",
"SemanticTextSplitter",
"PriorityTextSplitter",
"ModelConfig",
"ParameterLoader", "ParameterLoader",
"TrainConfig",
"DatasetLoader",
"BpeTokenizer",
"TextGenerator", "TextGenerator",
"ChatGenerator", "ChatGenerator",
"StreamGenerator", "StreamGenerator",
@@ -42,15 +55,7 @@ __all__ = [
"RetrievalGenerator", "RetrievalGenerator",
"EmbeddingEncoder", "EmbeddingEncoder",
# trainer
"Trainer", "Trainer",
"DatasetLoader",
"TrainConfig",
"StrategyFactory", "StrategyFactory",
"SchedulerFactory", "SchedulerFactory"
]
# utils
"Retriever",
"SemanticTextSplitter",
"PriorityTextSplitter",
]
+2 -2
View File
@@ -1,7 +1,7 @@
from torch import Tensor from torch import Tensor
from typing import List, Tuple, Generator, Union from typing import List, Tuple, Generator, Union
from khaosz.core.generator import ( from khaosz.inference.generator import (
TextGenerator, TextGenerator,
ChatGenerator, ChatGenerator,
StreamGenerator, StreamGenerator,
@@ -9,7 +9,7 @@ from khaosz.core.generator import (
RetrievalGenerator, RetrievalGenerator,
EmbeddingEncoder EmbeddingEncoder
) )
from khaosz.core.parameter import ParameterLoader from khaosz.config.param_config import ParameterLoader
class Khaosz: class Khaosz:
+18
View File
@@ -0,0 +1,18 @@
from khaosz.config.model_config import ModelConfig
from khaosz.config.param_config import BaseModelIO, ModelParameter, Checkpoint, ParameterLoader
from khaosz.config.schedule_config import ScheduleConfig, CosineScheduleConfig, SGDRScheduleConfig
from khaosz.config.train_config import TrainConfig
__all__ = [
"BaseModelIO",
"ModelParameter",
"Checkpoint",
"ParameterLoader",
"ModelConfig",
"TrainConfig",
"ScheduleConfig",
"CosineScheduleConfig",
"SGDRScheduleConfig",
]
+37
View File
@@ -0,0 +1,37 @@
import json
from dataclasses import asdict, dataclass
from typing import Any, Dict, Optional, Self
@dataclass
class ModelConfig:
# basic config
vocab_size: Optional[int] = None
n_dim: Optional[int] = None
n_head: Optional[int] = None
n_layer: Optional[int] = None
m_len: Optional[int] = None
norm_eps: Optional[float] = None
d_ffn: Optional[int] = None
tie_weight: Optional[bool] = None
# GQA
n_kvhead: Optional[int] = None
def load(self, config_path: str) -> Self:
with open(config_path, 'r') as f:
config: Dict[str, Any] = json.load(f)
for key, value in config.items():
if hasattr(self, key):
setattr(self, key, value)
return self
def save(self, config_path: str) -> None:
config_dict = asdict(self)
config_dict = {k: v for k, v in config_dict.items() if v is not None}
with open(config_path, 'w') as f:
json.dump(config_dict, f, indent=4)
@@ -8,8 +8,9 @@ from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Self, Union from typing import Any, Dict, List, Optional, Self, Union
from pathlib import Path from pathlib import Path
from khaosz.core.tokenizer import BpeTokenizer from khaosz.data.tokenizer import BpeTokenizer
from khaosz.core.transformer import TransformerConfig, Transformer from khaosz.config.model_config import ModelConfig
from khaosz.model.transformer import Transformer
class BaseModelIO: class BaseModelIO:
@@ -19,11 +20,11 @@ class BaseModelIO:
self, self,
model: Optional[nn.Module] = None, model: Optional[nn.Module] = None,
tokenizer: Optional[BpeTokenizer] = None, tokenizer: Optional[BpeTokenizer] = None,
config: Optional[TransformerConfig] = None config: Optional[ModelConfig] = None
): ):
self.model = model self.model = model
self.tokenizer = tokenizer or BpeTokenizer() self.tokenizer = tokenizer or BpeTokenizer()
self.config = config or TransformerConfig() self.config = config or ModelConfig()
def _get_file_paths(self, directory: Union[str, Path]) -> dict[str, Path]: def _get_file_paths(self, directory: Union[str, Path]) -> dict[str, Path]:
"""Get standardized file paths for model components.""" """Get standardized file paths for model components."""
@@ -78,8 +79,8 @@ class ModelParameter(BaseModelIO):
default_factory=BpeTokenizer, default_factory=BpeTokenizer,
metadata={"help": "Tokenizer for the model."} metadata={"help": "Tokenizer for the model."}
) )
config: TransformerConfig = field( config: ModelConfig = field(
default_factory=TransformerConfig, default_factory=ModelConfig,
metadata={"help": "Transformer model configuration."} metadata={"help": "Transformer model configuration."}
) )
@@ -102,15 +103,15 @@ class Checkpoint(BaseModelIO):
default_factory=BpeTokenizer, default_factory=BpeTokenizer,
metadata={"help": "Tokenizer for the model."} metadata={"help": "Tokenizer for the model."}
) )
config: TransformerConfig = field( config: ModelConfig = field(
default_factory=TransformerConfig, default_factory=ModelConfig,
metadata={"help": "Transformer model configuration."} metadata={"help": "Transformer model configuration."}
) )
optim_state: Dict[str, Any] = field( optimizer_state: Dict[str, Any] = field(
default=None, default=None,
metadata={"help": "Optimizer state."} metadata={"help": "Optimizer state."}
) )
sampler_state: Dict[str, Any] = field( scheduler_state: Dict[str, Any] = field(
default=None, default=None,
metadata={"help": "Sampler state."} metadata={"help": "Sampler state."}
) )
@@ -118,13 +119,21 @@ class Checkpoint(BaseModelIO):
default_factory=list, default_factory=list,
metadata={"help": "List of training losses."} metadata={"help": "List of training losses."}
) )
epoch: int = field(
default=0,
metadata={"help": "Current epoch."}
)
batch_iter: int = field(
default=0,
metadata={"help": "Current iteration."}
)
def _get_training_paths(self, directory: Union[str, Path]) -> dict[str, Path]: def _get_training_paths(self, directory: Union[str, Path]) -> dict[str, Path]:
paths = self._get_file_paths(directory) paths = self._get_file_paths(directory)
paths.update({ paths.update({
"loss_list": paths["model"].parent / "loss.pkl", "loss_list": paths["model"].parent / "loss.pkl",
"loss_plot": paths["model"].parent / "loss.png", "loss_plot": paths["model"].parent / "loss.png",
"optim_state": paths["model"].parent / "optim_state.pkl", "optimizer_state": paths["model"].parent / "optimizer_state.pkl",
"sampler_state": paths["model"].parent / "sampler_state.pkl" "sampler_state": paths["model"].parent / "sampler_state.pkl"
}) })
return paths return paths
@@ -140,12 +149,12 @@ class Checkpoint(BaseModelIO):
pkl.dump(self.loss_list, f) pkl.dump(self.loss_list, f)
# Save optimizer state # Save optimizer state
with open(str(paths["optim_state"]), "wb") as f: with open(str(paths["optimizer_state"]), "wb") as f:
pkl.dump(self.optim_state, f) pkl.dump(self.optimizer_state, f)
# Save sampler state # Save sampler state
with open(str(paths["sampler_state"]), "wb") as f: with open(str(paths["sampler_state"]), "wb") as f:
pkl.dump(self.sampler_state, f) pkl.dump(self.scheduler_state, f)
def load_training_state(self, load_dir: Union[str, Path]) -> Self: def load_training_state(self, load_dir: Union[str, Path]) -> Self:
paths = self._get_training_paths(load_dir) paths = self._get_training_paths(load_dir)
@@ -156,14 +165,14 @@ class Checkpoint(BaseModelIO):
self.loss_list = pkl.load(f) self.loss_list = pkl.load(f)
# Load optimizer state # Load optimizer state
if paths["optim_state"].exists(): if paths["optimizer_state"].exists():
with open(str(paths["optim_state"]), "rb") as f: with open(str(paths["optimizer_state"]), "rb") as f:
self.optim_state = pkl.load(f) self.optimizer_state = pkl.load(f)
# Load sampler state # Load sampler state
if paths["sampler_state"].exists(): if paths["sampler_state"].exists():
with open(str(paths["sampler_state"]), "rb") as f: with open(str(paths["sampler_state"]), "rb") as f:
self.sampler_state = pkl.load(f) self.scheduler_state = pkl.load(f)
return self return self
@@ -172,11 +181,11 @@ class Checkpoint(BaseModelIO):
if not self.loss_list: if not self.loss_list:
return return
current_iter = len(self.loss_list) batch_iter = len(self.loss_list)
plt.figure(figsize=(10, 6)) plt.figure(figsize=(10, 6))
plt.plot(self.loss_list) plt.plot(self.loss_list)
plt.title(f"Training Loss - Iteration {current_iter}") plt.title(f"Training Loss - Iteration {batch_iter}")
plt.xlabel("Batch") plt.xlabel("Batch")
plt.ylabel("Loss") plt.ylabel("Loss")
plt.grid(True) plt.grid(True)
@@ -221,7 +230,7 @@ class ParameterLoader:
def create_checkpoint( def create_checkpoint(
model: nn.Module, model: nn.Module,
tokenizer: BpeTokenizer, tokenizer: BpeTokenizer,
config: TransformerConfig, config: ModelConfig,
loss_list: Optional[list[float]] = None, loss_list: Optional[list[float]] = None,
optimizer: Optional[optim.Optimizer] = None, optimizer: Optional[optim.Optimizer] = None,
) -> Checkpoint: ) -> Checkpoint:
@@ -232,6 +241,4 @@ class ParameterLoader:
config=config, config=config,
loss_list=loss_list or [], loss_list=loss_list or [],
optimizer_state=optimizer optimizer_state=optimizer
) )
+87
View File
@@ -0,0 +1,87 @@
from typing import Any, Literal, Dict
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
@dataclass
class ScheduleConfig(ABC):
schedule_type: str = field(
default="cosine",
metadata={
"help": "Type of learning rate schedule.",
"choices": ["cosine", "sgdr"]
}
)
warmup_steps: int = field(
default=1000,
metadata={"help": "Number of warmup steps."}
)
min_rate: float = field(
default=0.05,
metadata={"help": "Minimum learning rate multiplier."}
)
@abstractmethod
def get_kwargs(self) -> Dict[str, Any]:
raise NotImplementedError
def validate(self) -> None:
"""Validate configuration parameters."""
if self.warmup_steps < 0:
raise ValueError(f"warmup_steps must be non-negative, got {self.warmup_steps}")
if not 0 <= self.min_rate <= 1:
raise ValueError(f"min_rate must be between 0 and 1, got {self.min_rate}")
@dataclass
class CosineScheduleConfig(ScheduleConfig):
total_steps: int = field(
default=None,
metadata={"help": "Total training steps for cosine schedule."}
)
schedule_type: Literal["cosine"] = "cosine"
def get_kwargs(self) -> Dict[str, Any]:
if self.total_steps is None:
raise ValueError("total_steps must be specified for cosine schedule")
return {
"schedule_type": self.schedule_type,
"warmup_steps": self.warmup_steps,
"lr_decay_steps": self.total_steps - self.warmup_steps,
"min_rate": self.min_rate
}
def validate(self) -> None:
super().validate()
if self.total_steps is not None and self.total_steps <= self.warmup_steps:
raise ValueError(f"total_steps ({self.total_steps}) must be greater than warmup_steps ({self.warmup_steps})")
@dataclass
class SGDRScheduleConfig(ScheduleConfig):
cycle_length: int = field(
default=1000,
metadata={"help": "Length of the first cycle in steps."}
)
t_mult: int = field(
default=2,
metadata={"help": "Multiplier for cycle length growth."}
)
schedule_type: Literal["sgdr"] = "sgdr"
def get_kwargs(self) -> Dict[str, Any]:
return {
"schedule_type": self.schedule_type,
"warmup_steps": self.warmup_steps,
"cycle_length": self.cycle_length,
"min_rate": self.min_rate,
"t_mult": self.t_mult
}
def validate(self) -> None:
super().validate()
if self.cycle_length <= 0:
raise ValueError(f"cycle_length must be positive, got {self.cycle_length}")
if self.t_mult < 1:
raise ValueError(f"t_mult must be >= 1, got {self.t_mult}")
+72
View File
@@ -0,0 +1,72 @@
from dataclasses import dataclass, field
from typing import Optional, TYPE_CHECKING
from torch.utils.data import Dataset
from torch.optim import Optimizer
if TYPE_CHECKING:
from khaosz.trainer.strategy import BaseStrategy
@dataclass
class TrainConfig:
strategy: "BaseStrategy" = field(
default=None,
metadata={"help": "Training strategy."}
)
dataset: Dataset = field(
default=None,
metadata={"help": "Dataset for training."}
)
optimizer: Optimizer = field(
default=None,
metadata={"help": "Optimizer for training."}
)
checkpoint_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."}
)
start_epoch: int = field(
default=0,
metadata={"help": "Start epoch for training."}
)
start_batch: int = field(
default=0,
metadata={"help": "Start batch iteration for training."}
)
checkpoint_interval: int = field(
default=5000,
metadata={"help": "Number of iterations between checkpoints."}
)
accumulation_steps: 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."}
)
num_workers: int = field(
default=0,
metadata={"help": "Number of workers for dataloader."}
)
prefetch_factor: Optional[int] = field(
default=None,
metadata={"help": "Prefetch factor for dataloader."}
)
pin_memory: bool = field(
default=False,
metadata={"help": "Pin memory for dataloader."}
)
-27
View File
@@ -1,27 +0,0 @@
from khaosz.core.tokenizer import BpeTokenizer
from khaosz.core.transformer import Transformer, TransformerConfig
from khaosz.core.parameter import ParameterLoader, ModelParameter, Checkpoint
from khaosz.core.generator import (
TextGenerator,
ChatGenerator,
StreamGenerator,
BatchGenerator,
RetrievalGenerator,
EmbeddingEncoder
)
__all__ = [
"Transformer",
"TransformerConfig",
"BpeTokenizer",
"ParameterLoader",
"ModelParameter",
"Checkpoint",
"TextGenerator",
"ChatGenerator",
"StreamGenerator",
"BatchGenerator",
"RetrievalGenerator",
"EmbeddingEncoder"
]
-568
View File
@@ -1,568 +0,0 @@
import torch
from torch import Tensor
from typing import List, Tuple, Union, Optional, Generator, Self
from khaosz.core.parameter import ModelParameter
def build_prompt(query: str, history: Optional[List[Tuple[str, str]]] = None) -> str:
"""
Build prompt for query and history
Args:
query(str): query string
history(Optional[List[Tuple[str, str]]]): history list of query and response
Returns:
str: prompt string
"""
prompt_parts = []
if history is None:
history = []
for his_query, his_response in history:
prompt_parts.append(f"<|user|> {his_query} <|system|> <bos>{his_response}<eos>")
if query is not None:
prompt_parts.append(f"<|user|> {query} <|system|> <bos>")
return "\n".join(prompt_parts)
def pad_sequence(ids_list: List[List[int]], max_ids_len: int, pad_id: int) -> List[List[int]]:
"""
Pad a list of sequences to a fixed length.
Args:
ids_list (List[List[int]]): A list of sequences.
max_ids_len (int): The maximum length of sequences.
pad_id (int): The id to pad sequences.
Returns:
List[List[int]]: A list of padded sequences.
"""
new_ids_list = []
for ids in ids_list:
pad_len = max_ids_len - len(ids)
padded_seq = [pad_id] * pad_len + ids
new_ids_list.append(padded_seq)
return new_ids_list
def apply_sampling_strategies(
logits: Tensor,
temperature: float,
top_k: int,
top_p: float,
filter_value: float = -float("inf")
) -> Tensor:
"""
Apply sampling strategies to the logits tensor.
Args:
logits (Tensor): The logits tensor.
temperature (float): The temperature parameter.
top_k (int): The top-k parameter.
top_p (float): The top-p parameter.
filter_value (float, optional): The filter value. Defaults to -float("inf").
Returns:
Tensor: The sampled logits tensor.
"""
if temperature != 1.0:
logits = logits / temperature
if top_k > 0:
top_k = min(top_k, logits.size(-1))
indices_to_remove = logits < torch.topk(logits, top_k, dim=-1)[0][..., -1, None]
logits[indices_to_remove] = filter_value
if top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
sorted_indices_to_remove = cumulative_probs > top_p
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
indices_to_remove = torch.zeros_like(logits, dtype=torch.bool)
indices_to_remove.scatter_(
dim=1,
index=sorted_indices,
src=sorted_indices_to_remove
)
logits[indices_to_remove] = filter_value
return logits
class KVCacheManager:
def __init__(
self,
num_layers: int,
batch_size: int,
max_len: int,
num_heads: int,
head_dim: int,
device: torch.device = "cuda",
dtype: torch.dtype = torch.bfloat16
):
self.num_layers = num_layers
self.batch_size = batch_size
self.max_len = max_len
self.num_heads = num_heads
self.head_dim = head_dim
self.device = device
self.dtype = dtype
self._kv_cache: List[Tuple[Tensor, Tensor]] = None
self._seq_mask: Tensor = None
self._initialize()
def _initialize(self):
self._kv_cache = []
for _ in range(self.num_layers):
k_cache = torch.zeros(
(self.batch_size, self.max_len, self.num_heads, self.head_dim),
device=self.device, dtype=self.dtype
)
v_cache = torch.zeros(
(self.batch_size, self.max_len, self.num_heads, self.head_dim),
device=self.device, dtype=self.dtype
)
self._kv_cache.append((k_cache, v_cache))
self._seq_mask = torch.ones(
(self.batch_size, self.max_len),
device=self.device, dtype=torch.bool
)
def update(self, active_mask: Tensor):
for i in range(self.num_layers):
k_cache, v_cache = self._kv_cache[i]
new_k_cache, new_v_cache = k_cache[active_mask], v_cache[active_mask]
self._kv_cache[i] = (new_k_cache, new_v_cache)
self._seq_mask = self._seq_mask[active_mask]
def reset(self, full_reset=False):
if full_reset:
self._kv_cache = None
self._seq_mask = None
else:
self._initialize()
def set_seq_mask(self, input_ids: Tensor, pad_id: int):
batch_size, seq_len = input_ids.shape
bool_mask = (input_ids != pad_id)
self._seq_mask[: batch_size, : seq_len] = bool_mask
def get_kvcache(self) -> List[Tuple[Tensor, Tensor]]:
return self._kv_cache
def get_seq_mask(self) -> Tensor:
return self._seq_mask
class GeneratorCore:
def __init__(self, parameter: ModelParameter):
self.model = parameter.model
self.tokenizer = parameter.tokenizer
self.config = parameter.config
def compute_logits(
self,
input_ids: Tensor,
attn_mask: Optional[Tensor] = None,
kv_caches: Optional[List[Tuple[Tensor, Tensor]]] = None,
start_pos: int = 0
) -> Tuple[Tensor, int]:
with torch.inference_mode():
outputs = self.model(input_ids, attn_mask, kv_caches, start_pos)
logits = outputs["logits"][:, -1, :]
cache_increase = input_ids.size(-1)
return logits, cache_increase
def to(self, *args, **kargs) -> Self:
self.model.to(*args, **kargs)
return self
class EmbeddingEncoderCore:
def __init__(self, parameter: ModelParameter):
self.model = parameter.model
self.tokenizer = parameter.tokenizer
self.config = parameter.config
def encode(self, sentence: Union[str, List[str]]) -> Union[Tensor, List[Tensor]]:
with_batch = isinstance(sentence, list)
ids = self.tokenizer.encode(sentence)
batch_ids = ids if with_batch else [ids]
max_model_len = self.config.m_len
all_fragments = []
fragment_origin_idx = []
for i, seq in enumerate(batch_ids):
if len(seq) > max_model_len:
fragments = [seq[j:j+max_model_len] for j in range(0, len(seq), max_model_len)]
all_fragments.extend(fragments)
fragment_origin_idx.extend([i] * len(fragments))
else:
all_fragments.append(seq)
fragment_origin_idx.append(i)
#if empty fragments
if not all_fragments or not ids:
return [] if with_batch else torch.tensor([])
device = next(self.model.parameters()).device
max_len = min(max(len(seq) for seq in all_fragments), max_model_len)
padded_ids = []
masks = []
for seq in all_fragments:
pad_len = max_len - len(seq)
padded_seq = seq + [self.tokenizer.pad_id] * pad_len
mask = [token_id != self.tokenizer.pad_id for token_id in padded_seq]
padded_ids.append(padded_seq)
masks.append(mask)
input_tensor = torch.tensor(padded_ids, device=device, dtype=torch.long)
seq_mask = torch.tensor(masks, device=device, dtype=torch.bool)
with torch.inference_mode():
outputs = self.model(input_tensor, seq_mask)["hidden_states"]
# [num_fragments, seq_len, hidden_size]
fragment_embs = torch.mul(outputs, seq_mask.unsqueeze(-1))
sentence_embs: List[Tensor] = []
for i in range(len(batch_ids)):
indices = [idx for idx, orig_idx in enumerate(fragment_origin_idx) if orig_idx == i]
if indices is not None:
sum_frags = torch.sum(fragment_embs[indices, :, :], dim=1) # [frags, hidden_size]
length = torch.sum(seq_mask[indices, :], dim=1).unsqueeze(1) # [frags, 1]
emb = torch.sum(sum_frags / length, dim=0) # [frags, hidden_size]
sentence_embs.append(emb.flatten())
if with_batch:
return [emb.flatten() for emb in sentence_embs]
else:
return sentence_embs[0].flatten()
def to(self, *args, **kargs) -> Self:
self.model.to(*args, **kargs)
return self
class TextGenerator(GeneratorCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def generate(
self,
query: str,
temperature: float,
top_k: int,
top_p: float,
) -> str:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
device = next(self.model.parameters()).device
cache_manager = KVCacheManager(
num_layers=self.config.n_layer,
batch_size=1,
max_len=self.config.m_len,
num_heads=self.config.n_kvhead,
head_dim=self.config.n_dim // self.config.n_head,
device=device,
)
ids = self.tokenizer.encode(query)
input_ids = torch.tensor([ids], device=device, dtype=torch.long)
start_cache_pos = len(ids)
cur_cache_pos = 0
self.model.eval()
while len(ids) < self.config.m_len:
kv_caches = cache_manager.get_kvcache()
logits, cache_increase = self.compute_logits(
input_ids,
kv_caches=kv_caches,
start_pos=cur_cache_pos
)
logits = apply_sampling_strategies(logits, temperature, top_k, top_p)
probs = torch.softmax(logits, dim=-1)
next_token_id = torch.multinomial(probs, num_samples=1)
input_ids = next_token_id
ids.append(next_token_id.item())
cur_cache_pos += cache_increase
if next_token_id.item() in self.tokenizer.stop_ids:
break
response = self.tokenizer.decode(ids[start_cache_pos:])
return response
class ChatGenerator(GeneratorCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def generate(
self,
query: str,
history: List[Tuple[str, str]],
temperature: float,
top_k: int,
top_p: float,
) -> str:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
if history is None:
history = []
device = next(self.model.parameters()).device
cache_manager = KVCacheManager(
num_layers=self.config.n_layer,
batch_size=1,
max_len=self.config.m_len,
num_heads=self.config.n_kvhead,
head_dim=self.config.n_dim // self.config.n_head,
device=device,
)
ids = self.tokenizer.encode(build_prompt(query, history))
input_ids = torch.tensor([ids], device=device, dtype=torch.long)
cpy_history = history.copy()
start_cache_pos = len(ids)
cur_cache_pos = 0
self.model.eval()
while len(ids) < self.config.m_len:
kv_caches = cache_manager.get_kvcache()
logits, cache_increase = self.compute_logits(
input_ids,
kv_caches=kv_caches,
start_pos=cur_cache_pos
)
logits = apply_sampling_strategies(logits, temperature, top_k, top_p)
probs = torch.softmax(logits, dim=-1)
next_token_id = torch.multinomial(probs, num_samples=1)
input_ids = next_token_id
ids.append(next_token_id.item())
cur_cache_pos += cache_increase
if next_token_id.item() in self.tokenizer.stop_ids:
break
response = self.tokenizer.decode(ids[start_cache_pos:])
cpy_history.append((query, response))
return response, cpy_history
class StreamGenerator(GeneratorCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def generate(
self,
query: str,
history: List[Tuple[str, str]],
temperature: float,
top_k: int,
top_p: float,
) -> Generator[Tuple[str, List[Tuple[str, str]]], None, None]:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
if history is None:
history = []
device = next(self.model.parameters()).device
cache_manager = KVCacheManager(
num_layers=self.config.n_layer,
batch_size=1,
max_len=self.config.m_len,
num_heads=self.config.n_kvhead,
head_dim=self.config.n_dim // self.config.n_head,
device=device,
)
ids = self.tokenizer.encode(build_prompt(query, history))
input_ids = torch.tensor([ids], device=device, dtype=torch.long)
cpy_history = history.copy()
start_cache_pos = len(ids)
cur_cache_pos = 0
self.model.eval()
while len(ids) < self.config.m_len:
kv_caches = cache_manager.get_kvcache()
logits, cache_increase = self.compute_logits(
input_ids,
kv_caches=kv_caches,
start_pos=cur_cache_pos
)
logits = apply_sampling_strategies(logits, temperature, top_k, top_p)
probs = torch.softmax(logits, dim=-1)
next_token_id = torch.multinomial(probs, num_samples=1)
input_ids = next_token_id
ids.append(next_token_id.item())
cur_cache_pos += cache_increase
response = self.tokenizer.decode(ids[start_cache_pos:])
yield response, cpy_history + [(query, response)]
if next_token_id.item() in self.tokenizer.stop_ids:
yield response + "\n", cpy_history + [(query, response)]
break
class BatchGenerator(GeneratorCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def generate(
self,
queries: List[str],
histories: List[List[Tuple[str, str]]],
temperature: float,
top_k: int,
top_p: float
) -> List[str]:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
batch_size = len(queries)
if histories is None:
histories = [[] for _ in range(batch_size)]
prompts = [build_prompt(query, history) for query, history in zip(queries, histories)]
ids_list = [self.tokenizer.encode(prompt) for prompt in prompts]
max_ids_len = max(len(ids) for ids in ids_list)
ids_list = pad_sequence(ids_list, max_ids_len, self.tokenizer.pad_id)
device = next(self.model.parameters()).device
cache_manager = KVCacheManager(
num_layers=self.config.n_layer,
batch_size=batch_size,
max_len=self.config.m_len,
num_heads=self.config.n_kvhead,
head_dim=self.config.n_dim // self.config.n_head,
device=device,
)
input_tensor = torch.tensor(ids_list, device=device, dtype=torch.long)
cache_manager.set_seq_mask(input_tensor, self.tokenizer.pad_id)
activate_task_mask = [True] * batch_size
start_cache_pos = max_ids_len
cur_cache_pos = 0
while max_ids_len < self.config.m_len and sum(activate_task_mask) != 0:
kv_caches = cache_manager.get_kvcache()
attn_mask =cache_manager.get_seq_mask()
logits, cache_increase = self.compute_logits(
input_tensor,
attn_mask=attn_mask,
kv_caches=kv_caches,
start_pos=cur_cache_pos
)
cur_cache_pos += cache_increase
logits = apply_sampling_strategies(logits, temperature, top_k, top_p)
probs = torch.softmax(logits, dim=-1)
next_token_id = torch.multinomial(probs, num_samples=1)
active_mask = []
c_ids = 0
for i in range(batch_size):
if activate_task_mask[i]:
token = next_token_id[c_ids, :].item()
ids_list[i].append(token)
c_ids += 1
is_active = not token in self.tokenizer.stop_ids
activate_task_mask[i] = is_active
active_mask.append(is_active)
active_mask = torch.tensor(active_mask, device=device, dtype=torch.bool)
cache_manager.update(active_mask)
input_tensor = next_token_id[active_mask, :]
max_ids_len += 1
responses = [str()] * batch_size
for i in range(batch_size):
responses[i] = self.tokenizer.decode(ids_list[i][start_cache_pos:])
histories[i].append((queries[i], responses[i]))
return responses
class RetrievalGenerator(GeneratorCore):
def __init__(self, retriever_parameter: ModelParameter):
super().__init__(retriever_parameter)
def generate(
self,
retrieved: List[str],
query: str,
history: List[Tuple[str, str]],
temperature: float,
top_k: int,
top_p: float,
) -> str:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
if history is None:
history = []
retrieved = "\n".join([f"{idx + 1}. {key}" for idx, key in enumerate(retrieved)]) if retrieved else ""
retrieved_query = f"{retrieved}<eos>\n\n根据以上内容回答: {query}" if retrieved else query
parameter = ModelParameter(self.model, self.tokenizer, self.config)
return ChatGenerator(parameter).generate(
retrieved_query,
history,
temperature=temperature,
top_k=top_k,
top_p=top_p,
)
class EmbeddingEncoder(EmbeddingEncoderCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def encode(self, sentence: Union[str, List[str]]) -> Union[Tensor, List[Tensor]]:
return super().encode(sentence)
-346
View File
@@ -1,346 +0,0 @@
import json
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from torch.nn import init
from dataclasses import asdict, dataclass
from typing import List, Optional, Self, Tuple
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
"""
Repeat k times along the dimension for attention heads.
Args:
x (Tensor): The input tensor.
n_rep (int): The number of repetitions.
Returns:
Tensor: The repeated tensor.
"""
bs, slen, n_heads, head_dim = x.shape
if n_rep == 1:
return x
return (
x[:, :, :, None, :]
.expand(bs, slen, n_heads, n_rep, head_dim)
.reshape(bs, slen, n_heads * n_rep, head_dim)
)
def get_rotary_emb(
dim: int,
max_len: int,
base: float = 10000,
device: torch.device = "cuda",
) -> torch.Tensor:
"""
Get the rotary embedding for the given dimension and maximum length.
Args:
dim (int): The dimension of the input.
max_len (int): The maximum length of the input.
base (float, optional): The base for the frequency. Defaults to 10000.
device (torch.device, optional): The device to use. Defaults to "cuda".
Returns:
Tensor: The rotary embedding tensor.
"""
theta = base ** (-torch.arange(0, dim, 2, device=device).float() / dim)
t = torch.arange(0, max_len, device=device).float()
freqs = torch.outer(t, theta)
freqs_cis = torch.polar(torch.ones_like(freqs), freqs)
return freqs_cis
def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
"""
Apply rotary embedding to the input tensor.
Args:
x (Tensor): The input tensor.
freqs_cis (Tensor): The rotary embedding tensor.
Returns:
Tensor: The output tensor.
"""
dtype = x.dtype
seq_len = x.size(1)
x_complex = torch.view_as_complex(x.view(*x.shape[:-1], -1, 2).float())
freqs_cis = freqs_cis.reshape(1, seq_len, 1, -1)
x_out = torch.view_as_real(x_complex * freqs_cis).flatten(3)
return x_out.to(dtype)
def process_attention_mask(
seq_mask: Tensor,
start_pos: int = 0,
seq_len: int = 0,
is_causal: bool = False,
device: torch.device = "cuda",
dtype: torch.dtype = torch.float32
) -> Tensor:
"""
Create attention mask for GQA
Args:
seq_mask (Tensor): A tensor indicating whether each position is valid or not.
start_pos (int): The starting position of the sequence.
seq_len (int): The length of the sequence.
is_causal (bool): Whether the attention is causal or not.
device (torch.device): The device to use.
Returns:
Tensor: The attention mask tensor.
"""
if seq_mask is None:
if start_pos != 0:
# for single prompt chat
seq_mask = torch.ones((1, seq_len), dtype=torch.bool, device=device)
else:
return None
if seq_mask.dim() > 2:
# shape (bsz, seq_len) or (bsz,n_heads, seq_len, seq_len + start_pos)
# if ndim > 2, it's 4D tensor
return seq_mask
batch_size = seq_mask.size(0)
seq_mask = seq_mask[:, :start_pos + seq_len].to(device=device, dtype=torch.bool)
# (bsz, start_pos + seq_len)
expanded_mask = seq_mask.unsqueeze(1).expand(batch_size, seq_len, start_pos + seq_len)
# (bsz, seq_len, start_pos + seq_len)
if is_causal:
causal_mask = torch.tril(
torch.ones((seq_len, start_pos + seq_len), dtype=torch.bool, device=device),
diagonal=start_pos
)
causal_mask = causal_mask.unsqueeze(0).expand(batch_size, seq_len, start_pos + seq_len)
expanded_mask = expanded_mask & causal_mask
attention_mask = torch.zeros_like(expanded_mask, dtype=dtype, device=device)
attention_mask = attention_mask.masked_fill_(~expanded_mask, -torch.finfo(dtype).max / 2).unsqueeze(1)
# (bsz, 1, seq_len, seq_len + start_pos)
return attention_mask
@dataclass
class TransformerConfig:
# basic config
vocab_size: Optional[int] = None
n_dim: Optional[int] = None
n_head: Optional[int] = None
n_layer: Optional[int] = None
m_len: Optional[int] = None
norm_eps: Optional[float] = None
d_ffn: Optional[int] = None
# GQA
n_kvhead: Optional[int] = None
def load(self, config_path: str) -> Self:
with open(config_path, 'r') as f:
config: dict = json.load(f)
for key, value in config.items():
if hasattr(self, key):
setattr(self, key, value)
return self
def save(self, config_path: str) -> None:
config_dict = asdict(self)
config_dict = {k: v for k, v in config_dict.items() if v is not None}
with open(config_path, 'w') as f:
json.dump(config_dict, f, indent=4)
class Linear(nn.Module):
def __init__(self, in_dim: int, out_dim: int, bias: bool=False):
super().__init__()
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
init.normal_(self.weight, mean=0, std=0.006)
def forward(self, x: Tensor) -> Tensor:
return F.linear(x, self.weight, self.bias)
class RMSNorm(nn.Module):
def __init__(self, n_dim, norm_eps):
super().__init__()
self.weight = nn.Parameter(torch.ones(n_dim))
self.norm_eps = norm_eps
def forward(self, x: Tensor) -> Tensor:
dtype = x.dtype
x = x.float()
mean_square = torch.mean(torch.pow(x, 2), dim=-1, keepdim=True)
norm = x * torch.rsqrt(mean_square + self.norm_eps)
norm = norm.to(dtype)
out = norm * self.weight
return out
class MLP(nn.Module):
def __init__(self, n_dim: int, d_ffn: int):
super().__init__()
self.up = Linear(n_dim, d_ffn)
self.gate = Linear(n_dim, d_ffn)
self.down = Linear(d_ffn, n_dim)
def forward(self, x: Tensor) -> Tensor:
gated = self.up(x) * F.silu(self.gate(x))
out = self.down(gated)
return out
class GQA(nn.Module):
def __init__(
self,
n_dim: int,
n_head: int,
n_kvhead: int,
):
super().__init__()
assert n_dim % n_head == 0
assert n_head % n_kvhead == 0
self.head_dim = n_dim // n_head
self.n_dim = n_dim
self.n_heads = n_head
self.n_kvheads = n_kvhead
self.n_rep = n_head // n_kvhead
self.q_proj = Linear(n_dim, n_head * self.head_dim)
self.k_proj = Linear(n_dim, n_kvhead * self.head_dim)
self.v_proj = Linear(n_dim, n_kvhead * self.head_dim)
self.o_proj = Linear(n_dim, n_dim)
def forward(
self,
x: Tensor,
freqs_cis: Tensor,
mask: Tensor = None,
kv_cache: Optional[Tuple[Tensor, Tensor]] = None,
start_pos: int = 0
) -> Tensor:
bsz, seq_len, _ = x.size()
# x(bsz, seq_len, n_heads * head_dim) -> (bsz, seq_len, n_heads, head_dim)
q = self._split_heads(self.q_proj(x), self.n_heads)
k = self._split_heads(self.k_proj(x), self.n_kvheads)
v = self._split_heads(self.v_proj(x), self.n_kvheads)
q, k = apply_rotary_emb(q, freqs_cis), apply_rotary_emb(k, freqs_cis)
if kv_cache is not None:
k_cache, v_cache = kv_cache
# copy to cache
k_cache[:bsz, start_pos:start_pos + seq_len] = k
v_cache[:bsz, start_pos:start_pos + seq_len] = v
# get cache
k = k_cache[:bsz, :start_pos + seq_len]
v = v_cache[:bsz, :start_pos + seq_len]
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
# (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim)
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
sdqa_out = F.scaled_dot_product_attention(q, k, v, mask, is_causal=(mask == None)).permute(0, 2, 1, 3)
out = self.o_proj(sdqa_out.contiguous().view(bsz, seq_len, -1))
return out
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
batch_size, seq_len, _ = x.shape
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
return x
class DecoderBlock(nn.Module):
def __init__(self, n_dim, n_head, d_ffn, n_kvhead, norm_eps):
super().__init__()
self.attention = GQA(n_dim, n_head, n_kvhead)
self.norm_attn = RMSNorm(n_dim, norm_eps)
self.ffn = MLP(n_dim, d_ffn)
self.norm_ffn = RMSNorm(n_dim, norm_eps)
def forward(
self,
x: Tensor,
freqs_cis: Tensor,
attention_mask: Optional[Tensor] = None,
kv_cache: Optional[Tuple[Tensor, Tensor]] = None,
start_pos: int = 0
) -> Tensor:
# attention
attn_output = self.attention(
self.norm_attn(x),
freqs_cis,
attention_mask,
kv_cache,
start_pos
)
x = attn_output + x
# feed forward
x = self.ffn(self.norm_ffn(x)) + x
return x
class Transformer(nn.Module):
def __init__(self, config: TransformerConfig):
super().__init__()
self.embedding = nn.Parameter(torch.empty(config.vocab_size, config.n_dim))
self.layers = nn.ModuleList([
DecoderBlock(
config.n_dim,
config.n_head,
config.d_ffn,
config.n_kvhead,
config.norm_eps
)
for _ in range(config.n_layer)
])
self.norm = RMSNorm(config.n_dim, config.norm_eps)
self.freq_cis = get_rotary_emb(config.n_dim // config.n_head, config.m_len)
init.normal_(self.embedding, mean=0, std=0.02)
def forward(
self,
input_ids: Tensor,
input_mask: Optional[Tensor]=None,
persistent_key_values: Optional[List[Tuple[Tensor, Tensor]]]=None,
start_pos: int = 0
) -> Tensor:
assert input_ids.ndim == 2
seq_len = input_ids.size(-1)
x = F.embedding(input_ids, self.embedding)
self.freq_cis = self.freq_cis.to(x.device)
freqs_cis = self.freq_cis[start_pos:start_pos+seq_len]
has_kvcache = persistent_key_values is not None
attn_mask = process_attention_mask(
input_mask,
start_pos=start_pos,
seq_len=seq_len,
is_causal=has_kvcache,
device=x.device,
dtype=x.dtype
)
for i, layer in enumerate(self.layers):
kv_cache = persistent_key_values[i] if persistent_key_values else None
x = layer(x, freqs_cis, attn_mask, kv_cache, start_pos)
hidden_states = self.norm(x)
logits = F.linear(hidden_states, self.embedding)
return {
"logits": logits,
"hidden_states": hidden_states
}
+26
View File
@@ -0,0 +1,26 @@
from khaosz.data.data_util import (
BaseDataset,
SeqDataset,
DpoDataset,
SftDataset,
PpoDataset,
MutiSegmentFetcher,
ResumeableRandomSampler,
DatasetLoader,
load_pkl_files,
)
from khaosz.data.tokenizer import BpeTokenizer
__all__ = [
"BaseDataset",
"SeqDataset",
"DpoDataset",
"SftDataset",
"PpoDataset",
"MutiSegmentFetcher",
"ResumeableRandomSampler",
"DatasetLoader",
"load_pkl_files",
"BpeTokenizer"
]
@@ -4,8 +4,8 @@ import pickle as pkl
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from torch import Tensor from torch import Tensor
from torch.utils.data import Dataset, Sampler from torch.utils.data import Dataset, Sampler
from typing import Callable, List, Dict, Literal, Union from typing import Callable, List, Dict, Literal, Optional, Union
MutiSeg = Dict[str, List[Tensor]] MutiSeg = Dict[str, List[Tensor]]
Seg = Dict[str, Tensor] Seg = Dict[str, Tensor]
@@ -25,36 +25,6 @@ def load_pkl_files(paths: List[str]):
return segments, total_samples return segments, total_samples
def build_attention_mask(input_ids: Tensor, user_token_id: int, multi_turn: bool) -> Tensor:
seq_len = input_ids.size(0)
turn_id = input_ids.eq(user_token_id).cumsum(dim=-1)
iq = turn_id.view(seq_len, 1)
ik = turn_id.view(1, seq_len)
# fix the causual attention mask(iq >= ik condition)
seq_mask = (iq >= ik) if multi_turn else (iq == ik)
attention_mask = torch.tril(seq_mask)
# fix the shape (bsz, 1, seq_len, seq_len) unsqueeze for broadcast
return attention_mask.unsqueeze(0)
def build_loss_mask(input_ids: Tensor, bos_token_id: int, eos_token_id: int) -> Tensor:
token_markers = torch.zeros_like(input_ids, dtype=torch.int8)
is_bos_token = input_ids.eq(bos_token_id)
is_eos_token = input_ids.eq(eos_token_id)
# fix the eos_token_id bug(change target_ids to input_ids)
token_markers[is_bos_token] = 1
token_markers[is_eos_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 BaseSegmentFetcher: class BaseSegmentFetcher:
def __init__(self, segments: List[Tensor]): def __init__(self, segments: List[Tensor]):
@@ -72,8 +42,9 @@ class BaseSegmentFetcher:
if begin_idx >= end_idx: if begin_idx >= end_idx:
return torch.tensor([], dtype=torch.long) return torch.tensor([], dtype=torch.long)
seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx - 1) # fix the range index bug
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx - 1) seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx)
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx)
result_segments = [] result_segments = []
@@ -110,12 +81,12 @@ class MutiSegmentFetcher:
class BaseDataset(Dataset, ABC): class BaseDataset(Dataset, ABC):
def __init__(self, chunk_size: int, device: str): def __init__(self, window_size: int, stride: int):
super().__init__() super().__init__()
self.segments: MutiSeg = {} self.segments: MutiSeg = {}
self.chunk_size = chunk_size self.window_size = window_size
self.total_samples = 0 self.stride = stride
self.device = device self.total_samples = None
def save(self, save_path: str): def save(self, save_path: str):
keys = list(self.segments.keys()) keys = list(self.segments.keys())
@@ -128,117 +99,99 @@ class BaseDataset(Dataset, ABC):
for i in range(segment_size): for i in range(segment_size):
formated_segment = {key: self.segments[key][i] for key in keys} formated_segment = {key: self.segments[key][i] for key in keys}
pkl.dump(formated_segment, open(f"{save_path}_{i}.pkl", "wb")) pkl.dump(formated_segment, open(f"{save_path}_{i}.pkl", "wb"))
def load(self, load_path: Union[str, List[str]]): def load(self, load_path: Union[str, List[str]]):
paths = [load_path] if isinstance(load_path, str) else load_path paths = [load_path] if isinstance(load_path, str) else load_path
self.segments, self.total_samples = load_pkl_files(paths) self.segments, self.total_samples = load_pkl_files(paths)
self.fetcher = MutiSegmentFetcher(self.segments) self.fetcher = MutiSegmentFetcher(self.segments)
def get_index(self, index: int) -> int:
begin_idx = min(index * self.stride, self.total_samples - self.window_size - 1)
end_idx = begin_idx + self.window_size
return begin_idx, end_idx
@abstractmethod @abstractmethod
def __getitem__(self, index: int) -> Dict[str, Tensor]: def __getitem__(self, index: int) -> Dict[str, Tensor]:
raise NotImplementedError raise NotImplementedError
def __len__(self) -> int: def __len__(self) -> int:
assert self.total_samples // self.chunk_size > 0 assert self.total_samples is not None
return self.total_samples // self.chunk_size if self.total_samples <= self.window_size:
return 0
return self.total_samples // self.stride + 1
class SeqDataset(BaseDataset): class SeqDataset(BaseDataset):
def __init__( def __init__(self, window_size: int, stride: int):
self, super().__init__(window_size, stride)
chunk_size,
device='cuda'
):
super().__init__(chunk_size, device)
self.fetcher = MutiSegmentFetcher(self.segments) self.fetcher = MutiSegmentFetcher(self.segments)
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor: def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
return self.fetcher.key_fetch(begin_idx, end_idx, "sequence") return self.fetcher.key_fetch(begin_idx, end_idx, "sequence")
def __getitem__(self, index): def __getitem__(self, index):
begin_idx = index * self.chunk_size # fix the range index bug
end_idx = min(begin_idx + self.chunk_size, self.total_samples - 1) begin_idx, end_idx = self.get_index(index)
x = self._fetch_data(begin_idx, end_idx).to(device=self.device, dtype=torch.long) x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(device=self.device, dtype=torch.long) y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
return {"input_ids": x, "target_ids": y} return {"input_ids": x, "target_ids": y}
class SftDataset(BaseDataset): class SftDataset(BaseDataset):
def __init__( def __init__(self, window_size: int, stride: int):
self, super().__init__(window_size, stride)
chunk_size,
bos_token_id,
eos_token_id,
user_token_id,
multi_turn=False,
device='cuda'
):
super().__init__(chunk_size, device)
self.fetcher = MutiSegmentFetcher(self.segments) self.fetcher = MutiSegmentFetcher(self.segments)
self.bos_token_id = bos_token_id
self.eos_token_id = eos_token_id
self.user_token_id = user_token_id
self.multi_turn = multi_turn
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor: def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
return self.fetcher.key_fetch(begin_idx, end_idx, key) return self.fetcher.key_fetch(begin_idx, end_idx, key)
def __getitem__(self, index): def __getitem__(self, index):
begin_idx = index * self.chunk_size begin_idx, end_idx = self.get_index(index)
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) x = self._fetch_data(begin_idx, end_idx, "sequence").to(dtype=torch.long)
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(device=self.device, dtype=torch.long) y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(dtype=torch.long)
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask").to(dtype=torch.bool)
# fix the eos_token_id bug(change target_ids to input_ids) return {"input_ids": x, "target_ids": y, "loss_mask": loss_mask}
loss_mask = build_loss_mask(x, self.bos_token_id, self.eos_token_id)
attn_mask = build_attention_mask(x, self.user_token_id, self.multi_turn)
return {"input_ids": x, "target_ids": y, "loss_mask": loss_mask, "attn_mask": attn_mask}
class DpoDataset(BaseDataset): class DpoDataset(BaseDataset):
def __init__(self, chunk_size: int, device="cuda"): def __init__(self, window_size: int, stride: int):
super().__init__(chunk_size, device) super().__init__(window_size, stride)
self.fetcher = MutiSegmentFetcher(self.segments) self.fetcher = MutiSegmentFetcher(self.segments)
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor: def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
return self.fetcher.key_fetch(begin_idx, end_idx, key) return self.fetcher.key_fetch(begin_idx, end_idx, key)
def __getitem__(self, index: int): def __getitem__(self, index: int):
start_idx = index * self.chunk_size begin_idx, end_idx = self.get_index(index)
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) chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
rejected = self._fetch_data(start_idx, end_idx, "rejected").to(device=self.device, dtype=torch.long) rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
chosen_mask = self._fetch_data(start_idx, end_idx, "chosen_mask").to(device=self.device, dtype=torch.bool) chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(dtype=torch.bool)
rejected_mask = self._fetch_data(start_idx, end_idx, "rejected_mask").to(device=self.device, dtype=torch.bool) rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(dtype=torch.bool)
return {"chosen": chosen, "rejected": rejected, "chosen_mask": chosen_mask, "rejected_mask": rejected_mask} return {"chosen": chosen, "rejected": rejected, "chosen_mask": chosen_mask, "rejected_mask": rejected_mask}
class PpoDataset(BaseDataset): class PpoDataset(BaseDataset):
def __init__(self, chunk_size: int, device="cuda"): def __init__(self, window_size: int, stride: int):
super().__init__(chunk_size, device) super().__init__(window_size, stride)
self.fetcher = MutiSegmentFetcher(self.segments) self.fetcher = MutiSegmentFetcher(self.segments)
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor: def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
return self.fetcher.key_fetch(begin_idx, end_idx, key) return self.fetcher.key_fetch(begin_idx, end_idx, key)
def __getitem__(self, index: int) -> Dict[str, Tensor]: def __getitem__(self, index: int) -> Dict[str, Tensor]:
begin_idx, end_idx = self.get_index(index)
begin_idx = index * self.chunk_size input_ids = self._fetch_data(begin_idx, end_idx, "input_ids"),
end_idx = min(begin_idx + self.chunk_size, self.total_samples - 1) actions = self._fetch_data(begin_idx, end_idx, "actions"),
logprobs = self._fetch_data(begin_idx, end_idx, "logprobs"),
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
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": input_ids, "actions": actions, "logprobs": logprobs, "rewards": rewards} return {"input_ids": input_ids, "actions": actions, "logprobs": logprobs, "rewards": rewards}
@@ -248,79 +201,56 @@ class DatasetLoader:
def load( def load(
train_type: Literal["seq", "sft", "dpo"], train_type: Literal["seq", "sft", "dpo"],
load_path: Union[str, List[str]], load_path: Union[str, List[str]],
max_len: int, window_size: int,
device: str, stride: Optional[int] = None,
**kwargs **kwargs
) -> BaseDataset: ) -> BaseDataset:
if stride is None:
stride = window_size
dataset_router: Dict[str, Callable[[int, torch.device], BaseDataset]] = { dataset_router: Dict[str, Callable[[int], BaseDataset]] = {
"seq": lambda m_len, device: SeqDataset(m_len, device=device), "seq": lambda window_size: SeqDataset(window_size, stride),
"sft": lambda m_len, device: SftDataset( "sft": lambda window_size: SftDataset(window_size, stride),
m_len, "dpo": lambda window_size: DpoDataset(window_size, stride),
device=device,
bos_token_id=kwargs.get("bos_token_id"),
eos_token_id=kwargs.get("eos_token_id"),
user_token_id=kwargs.get("user_token_id"),
multi_turn=kwargs.get("multi_turn")
),
"dpo": lambda m_len, device: DpoDataset(m_len, device=device),
} }
dataset = dataset_router[train_type](max_len, device) dataset = dataset_router[train_type](window_size)
dataset.load(load_path) dataset.load(load_path)
return dataset return dataset
class RandomSampler(Sampler[int]):
def __init__(self, data_source, generator=None, seed=42): class ResumeableRandomSampler(Sampler[int]):
self.data_source = data_source def __init__(self, data_source, start_epoch=0, start_iter=0, seed=42):
self.seed = seed self.num_samples = len(data_source)
self.epoch = 0 self.epoch = start_epoch
self.current_iter = 0 self.iter = start_iter
self._indices = None
if generator is None: generator = torch.Generator()
self.generator = torch.Generator() generator.manual_seed(seed)
self.generator.manual_seed(seed)
else: # consume previous epochs
self.generator = generator for _ in range(start_epoch):
torch.randperm(self.num_samples, generator=generator)
self.generator = generator
self._indices = None
def _generate_indices(self): def _get_indices(self):
n = len(self.data_source) current_epoch_indices = torch.randperm(self.num_samples, generator=self.generator).tolist()
self._indices = torch.randperm(n, generator=self.generator).tolist() self._indices = current_epoch_indices[self.iter % self.num_samples:]
def __iter__(self): def __iter__(self):
n = len(self.data_source)
if self._indices is None: if self._indices is None:
self._generate_indices() self._get_indices()
start = self.current_iter % n for i in self._indices:
for i in range(start, n): self.iter += 1
yield self._indices[i] yield i
self.current_iter += 1
self.epoch += 1 self.epoch += 1
self._indices = None self._indices = None
def __len__(self): def __len__(self):
return len(self.data_source) if self._indices is None:
self._get_indices()
def state_dict(self): return len(self._indices)
return {
'epoch': self.epoch,
'current_iter': self.current_iter,
'seed': self.seed,
'generator_state': self.generator.get_state() if self.generator else None,
'indices': self._indices
}
def load_state_dict(self, state_dict):
self.epoch = state_dict['epoch']
self.current_iter = state_dict['current_iter']
self.seed = state_dict['seed']
if self.generator and state_dict['generator_state'] is not None:
self.generator.set_state(state_dict['generator_state'])
self._indices = state_dict['indices']
@@ -8,7 +8,7 @@ from typing import List, Union
class BpeTokenizer: class BpeTokenizer:
def __init__(self, path=None): def __init__(self, path=None):
self._control_tokens = ["<bos>", "<eos>", "<pad>"] self._control_tokens = ["<bos>", "<eos>", "<pad>"]
self._special_tokens = ["<|user|>", "<|system|>"] self._special_tokens = ["<|im_start|>", "<|im_end|>"]
model = BPE() model = BPE()
tokenizer = Tokenizer(model) tokenizer = Tokenizer(model)
tokenizer.normalizer = normalizers.Sequence([ tokenizer.normalizer = normalizers.Sequence([
@@ -93,9 +93,8 @@ class BpeTokenizer:
@property @property
def stop_ids(self) -> List[int]: def stop_ids(self) -> List[int]:
stop_ids = [] stop_token = self._control_tokens + self._special_tokens
for token in self._control_tokens: stop_ids = [self._tokenizer.token_to_id(token) for token in stop_token]
stop_ids.append(self._tokenizer.token_to_id(token))
return stop_ids return stop_ids
@property @property
@@ -108,12 +107,4 @@ class BpeTokenizer:
@property @property
def pad_id(self) -> int: def pad_id(self) -> int:
return self._tokenizer.token_to_id("<pad>") return self._tokenizer.token_to_id("<pad>")
@property
def user_id(self) -> int:
return self._tokenizer.token_to_id("<|user|>")
@property
def system_id(self) -> int:
return self._tokenizer.token_to_id("<|system|>")
+240
View File
@@ -0,0 +1,240 @@
import torch
from torch import Tensor
from typing import Any, Callable, List, Tuple, Union, Optional, Self
from khaosz.config import ModelParameter, ModelConfig
def apply_sampling_strategies(
logits: Tensor,
temperature: float,
top_k: int,
top_p: float,
filter_value: float = -float("inf")
) -> Tensor:
"""
Apply sampling strategies to the logits tensor.
Args:
logits (Tensor): The logits tensor.
temperature (float): The temperature parameter.
top_k (int): The top-k parameter.
top_p (float): The top-p parameter.
filter_value (float, optional): The filter value. Defaults to -float("inf").
Returns:
Tensor: The sampled logits tensor.
"""
if temperature != 1.0:
logits = logits / temperature
if top_k > 0:
top_k = min(top_k, logits.size(-1))
indices_to_remove = logits < torch.topk(logits, top_k, dim=-1)[0][..., -1, None]
logits[indices_to_remove] = filter_value
if top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
sorted_indices_to_remove = cumulative_probs > top_p
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
indices_to_remove = torch.zeros_like(logits, dtype=torch.bool)
indices_to_remove.scatter_(
dim=1,
index=sorted_indices,
src=sorted_indices_to_remove
)
logits[indices_to_remove] = filter_value
return logits
class GeneratorCore:
def __init__(self, parameter: ModelParameter):
self.model = parameter.model
self.tokenizer = parameter.tokenizer
self.config = parameter.config
def generate_iterator(
self,
input_ids: Tensor,
temperature: float,
top_k: int,
top_p: float,
attn_mask: Optional[Tensor] = None,
kv_caches: Optional[List[Tuple[Tensor, Tensor]]] = None,
start_pos: int = 0
)-> Tuple[Tensor, int]:
with torch.inference_mode():
outputs = self.model(input_ids, attn_mask, kv_caches, start_pos)
logits = outputs["logits"][:, -1, :]
cache_increase = input_ids.size(-1)
logits = apply_sampling_strategies(logits, temperature, top_k, top_p)
probs = torch.softmax(logits, dim=-1)
next_token_id = torch.multinomial(probs, num_samples=1)
return next_token_id, cache_increase
def to(self, *args, **kargs) -> Self:
self.model.to(*args, **kargs)
return self
def generate_loop(
self,
input_ids: Tensor,
ids: List[int],
temperature: float,
top_k: int,
top_p: float,
attn_mask: Optional[Tensor] = None,
kv_caches: Optional[List[Tuple[Tensor, Tensor]]] = None,
start_pos: int = 0,
callback: Optional[Callable[..., Any]] = None
) -> List[int]:
cur_cache_pos = start_pos
for _ in range(len(ids), self.config.m_len):
next_token_id, cache_increase = self.generate_iterator(
input_ids, temperature, top_k, top_p, attn_mask, kv_caches, cur_cache_pos)
input_ids = next_token_id
ids.append(next_token_id.item())
cur_cache_pos += cache_increase
if callback:
callback(next_token_id.item(), ids.copy())
if next_token_id.item() in self.tokenizer.stop_ids:
break
return ids
class EmbeddingEncoderCore:
def __init__(self, parameter: ModelParameter):
self.model = parameter.model
self.tokenizer = parameter.tokenizer
self.config = parameter.config
def encode(self, sentence: Union[str, List[str]]) -> Union[Tensor, List[Tensor]]:
with_batch = isinstance(sentence, list)
ids = self.tokenizer.encode(sentence)
batch_ids = ids if with_batch else [ids]
max_model_len = self.config.m_len
all_fragments = []
fragment_origin_idx = []
for i, seq in enumerate(batch_ids):
if len(seq) > max_model_len:
fragments = [seq[j:j+max_model_len] for j in range(0, len(seq), max_model_len)]
all_fragments.extend(fragments)
fragment_origin_idx.extend([i] * len(fragments))
else:
all_fragments.append(seq)
fragment_origin_idx.append(i)
#if empty fragments
if not all_fragments or not ids:
return [] if with_batch else torch.tensor([])
device = next(self.model.parameters()).device
max_len = min(max(len(seq) for seq in all_fragments), max_model_len)
padded_ids = []
masks = []
for seq in all_fragments:
pad_len = max_len - len(seq)
padded_seq = seq + [self.tokenizer.pad_id] * pad_len
mask = [token_id != self.tokenizer.pad_id for token_id in padded_seq]
padded_ids.append(padded_seq)
masks.append(mask)
input_tensor = torch.tensor(padded_ids, device=device, dtype=torch.long)
seq_mask = torch.tensor(masks, device=device, dtype=torch.bool)
with torch.inference_mode():
outputs = self.model(input_tensor, seq_mask)["hidden_states"]
# [num_fragments, seq_len, hidden_size]
fragment_embs = torch.mul(outputs, seq_mask.unsqueeze(-1))
sentence_embs: List[Tensor] = []
for i in range(len(batch_ids)):
indices = [idx for idx, orig_idx in enumerate(fragment_origin_idx) if orig_idx == i]
if indices is not None:
sum_frags = torch.sum(fragment_embs[indices, :, :], dim=1) # [frags, hidden_size]
length = torch.sum(seq_mask[indices, :], dim=1).unsqueeze(1) # [frags, 1]
emb = torch.sum(sum_frags / length, dim=0) # [frags, hidden_size]
sentence_embs.append(emb.flatten())
if with_batch:
return [emb.flatten() for emb in sentence_embs]
else:
return sentence_embs[0].flatten()
def to(self, *args, **kargs) -> Self:
self.model.to(*args, **kargs)
return self
class KVCacheManager:
def __init__(
self,
config: ModelConfig,
batch_size: int,
device: torch.device = "cuda",
dtype: torch.dtype = torch.bfloat16
):
self.batch_size = batch_size
self.device = device
self.dtype = dtype
self.num_layers = config.n_layer
self.max_len = config.m_len
self.num_heads = config.n_kvhead
self.head_dim = config.n_dim //config.n_head
self._kv_cache: Tuple[Tensor, Tensor] = None
self._seq_mask: Tensor = None
self._initialize()
def _initialize(self):
k_cache = torch.zeros(
(self.batch_size, self.num_layers, self.max_len, self.num_heads, self.head_dim),
device=self.device, dtype=self.dtype
)
v_cache = torch.zeros(
(self.batch_size, self.num_layers, self.max_len, self.num_heads, self.head_dim),
device=self.device, dtype=self.dtype
)
self._kv_cache = (k_cache, v_cache)
self._seq_mask = torch.ones((self.batch_size, self.max_len), device=self.device, dtype=torch.bool)
def update(self, active_mask: Tensor):
k_cache, v_cache = self._kv_cache
self._kv_cache = (k_cache[active_mask], v_cache[active_mask])
self._seq_mask = self._seq_mask[active_mask]
def reset(self, full_reset=False):
if full_reset:
self._kv_cache = None
self._seq_mask = None
else:
self._initialize()
def set_seq_mask(self, input_ids: Tensor, pad_id: int):
batch_size, seq_len = input_ids.shape
bool_mask = (input_ids != pad_id)
self._seq_mask[: batch_size, : seq_len] = bool_mask
def get_kvcache(self) -> Tuple[Tensor, Tensor]:
return self._kv_cache
def get_seq_mask(self) -> Tensor:
return self._seq_mask
+296
View File
@@ -0,0 +1,296 @@
import torch
from torch import Tensor
from typing import List, Tuple, Union, Optional, Generator
from khaosz.inference.core import GeneratorCore, EmbeddingEncoderCore, KVCacheManager
from khaosz.config.param_config import ModelParameter
def build_prompt(
query: str,
init_prompt: Optional[str] = None,
history: Optional[List[Tuple[str, str]]] = None
) -> str:
"""
Build prompt in ChatML format for query and history
Args:
query(str): query string
history(Optional[List[Tuple[str, str]]]): history list of query and response
Returns:
str: prompt string in ChatML format
"""
prompt = f"<|im_start|>system\n{init_prompt}<|im_end|>\n" if init_prompt else ""
# (convert tuple format to ChatML)
if history:
for user_msg, assistant_msg in history:
prompt += f"<|im_start|>user\n{user_msg}<|im_end|>\n"
prompt += f"<|im_start|>assistant\n{assistant_msg}<|im_end|>\n"
prompt += f"<|im_start|>user\n{query}<|im_end|>\n"
prompt += "<|im_start|>assistant\n"
return prompt
def pad_sequence(ids_list: List[List[int]], max_ids_len: int, pad_id: int) -> List[List[int]]:
"""
Pad a list of sequences to a fixed length.
Args:
ids_list (List[List[int]]): A list of sequences.
max_ids_len (int): The maximum length of sequences.
pad_id (int): The id to pad sequences.
Returns:
List[List[int]]: A list of padded sequences.
"""
new_ids_list = []
for ids in ids_list:
pad_len = max_ids_len - len(ids)
padded_seq = [pad_id] * pad_len + ids
new_ids_list.append(padded_seq)
return new_ids_list
class TextGenerator(GeneratorCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def generate(
self,
query: str,
temperature: float,
top_k: int,
top_p: float,
) -> str:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
device = next(self.model.parameters()).device
cache_manager = KVCacheManager(self.config, 1, device=device)
ids = self.tokenizer.encode(query)
input_ids = torch.tensor([ids], device=device, dtype=torch.long)
start_cache_pos = len(ids)
cur_cache_pos = 0
self.model.eval()
kv_caches = cache_manager.get_kvcache()
ids = self.generate_loop(
input_ids, ids, temperature, top_k, top_p,
kv_caches=kv_caches,
start_pos=cur_cache_pos
)
response = self.tokenizer.decode(ids[start_cache_pos:])
return response
class ChatGenerator(GeneratorCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def generate(
self,
query: str,
history: List[Tuple[str, str]],
temperature: float,
top_k: int,
top_p: float,
) -> str:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
if history is None:
history = []
device = next(self.model.parameters()).device
cache_manager = KVCacheManager(self.config, 1, device=device)
ids = self.tokenizer.encode(build_prompt(query, history))
input_ids = torch.tensor([ids], device=device, dtype=torch.long)
start_cache_pos = len(ids)
cur_cache_pos = 0
self.model.eval()
kv_caches = cache_manager.get_kvcache()
ids = self.generate_loop(
input_ids, ids, temperature, top_k, top_p,
kv_caches=kv_caches,
start_pos=cur_cache_pos
)
response = self.tokenizer.decode(ids[start_cache_pos:])
return response
class StreamGenerator(GeneratorCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def generate(
self,
query: str,
history: List[Tuple[str, str]],
temperature: float,
top_k: int,
top_p: float,
) -> Generator[Tuple[str, List[Tuple[str, str]]], None, None]:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
if history is None:
history = []
device = next(self.model.parameters()).device
cache_manager = KVCacheManager(self.config, 1, device=device)
ids = self.tokenizer.encode(build_prompt(query, history))
input_ids = torch.tensor([ids], device=device, dtype=torch.long)
cpy_history = history.copy()
start_cache_pos = len(ids)
cur_cache_pos = 0
self.model.eval()
kv_caches = cache_manager.get_kvcache()
for _ in range(len(ids), self.config.m_len):
next_token_id, cache_increase = self.generate_iterator(
input_ids, temperature, top_k, top_p, kv_caches=kv_caches, start_pos=cur_cache_pos)
input_ids = next_token_id
ids.append(next_token_id.item())
cur_cache_pos += cache_increase
response = self.tokenizer.decode(ids[start_cache_pos:])
yield response, cpy_history + [(query, response)]
if next_token_id.item() in self.tokenizer.stop_ids:
yield response + "\n", cpy_history + [(query, response)]
break
class BatchGenerator(GeneratorCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def generate(
self,
queries: List[str],
histories: List[List[Tuple[str, str]]],
temperature: float,
top_k: int,
top_p: float
) -> List[str]:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
batch_size = len(queries)
if histories is None:
histories = [[] for _ in range(batch_size)]
prompts = [build_prompt(query, history) for query, history in zip(queries, histories)]
ids_list = [self.tokenizer.encode(prompt) for prompt in prompts]
max_ids_len = max(len(ids) for ids in ids_list)
ids_list = pad_sequence(ids_list, max_ids_len, self.tokenizer.pad_id)
device = next(self.model.parameters()).device
cache_manager = KVCacheManager(self.config, batch_size, device=device)
input_tensor = torch.tensor(ids_list, device=device, dtype=torch.long)
cache_manager.set_seq_mask(input_tensor, self.tokenizer.pad_id)
activate_task_mask = [True] * batch_size
start_cache_pos = max_ids_len
cur_cache_pos = 0
while max_ids_len < self.config.m_len and sum(activate_task_mask) != 0:
kv_caches = cache_manager.get_kvcache()
attn_mask =cache_manager.get_seq_mask()
next_token_id, cache_increase = self.generate_iterator(
input_tensor, temperature, top_k, top_p, attn_mask=attn_mask, kv_caches=kv_caches, start_pos=cur_cache_pos)
cur_cache_pos += cache_increase
active_mask = []
c_ids = 0
for i in range(batch_size):
if activate_task_mask[i]:
token = next_token_id[c_ids, :].item()
ids_list[i].append(token)
c_ids += 1
is_active = not token in self.tokenizer.stop_ids
activate_task_mask[i] = is_active
active_mask.append(is_active)
active_mask = torch.tensor(active_mask, device=device, dtype=torch.bool)
cache_manager.update(active_mask)
input_tensor = next_token_id[active_mask, :]
max_ids_len += 1
responses = [str()] * batch_size
for i in range(batch_size):
responses[i] = self.tokenizer.decode(ids_list[i][start_cache_pos:])
histories[i].append((queries[i], responses[i]))
return responses
class RetrievalGenerator(GeneratorCore):
def __init__(self, retriever_parameter: ModelParameter):
super().__init__(retriever_parameter)
def generate(
self,
retrieved: List[str],
query: str,
history: List[Tuple[str, str]],
temperature: float,
top_k: int,
top_p: float,
) -> str:
assert temperature >= 0.0
assert top_k >= 0
assert top_p >= 0.0 and top_p <= 1.0
if history is None:
history = []
retrieved = "\n".join([f"{idx + 1}. {key}" for idx, key in enumerate(retrieved)]) if retrieved else ""
retrieved_query = f"{retrieved}\n\n{query}" if retrieved else query
parameter = ModelParameter(self.model, self.tokenizer, self.config)
return ChatGenerator(parameter).generate(
retrieved_query,
history,
temperature=temperature,
top_k=top_k,
top_p=top_p,
)
class EmbeddingEncoder(EmbeddingEncoderCore):
def __init__(self, parameter: ModelParameter):
super().__init__(parameter)
def encode(self, sentence: Union[str, List[str]]) -> Union[Tensor, List[Tensor]]:
return super().encode(sentence)
+17
View File
@@ -0,0 +1,17 @@
from khaosz.model.module import (
Linear,
RMSNorm,
MLP,
GQA,
DecoderBlock,
)
from khaosz.model.transformer import Transformer
__all__ = [
"Linear",
"RMSNorm",
"MLP",
"GQA",
"DecoderBlock",
"Transformer"
]
+250
View File
@@ -0,0 +1,250 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from typing import Optional, Tuple
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
"""
Repeat k times along the dimension for attention heads.
Args:
x (Tensor): The input tensor.
n_rep (int): The number of repetitions.
Returns:
Tensor: The repeated tensor.
"""
bs, slen, n_heads, head_dim = x.shape
if n_rep == 1:
return x
return (
x[:, :, :, None, :]
.expand(bs, slen, n_heads, n_rep, head_dim)
.reshape(bs, slen, n_heads * n_rep, head_dim)
)
def get_rotary_emb(
dim: int,
max_len: int,
base: float = 10000,
) -> Tuple[Tensor, Tensor]:
"""
Get the rotary embedding for the given dimension and maximum length.
Args:
dim (int): The dimension of the input.
max_len (int): The maximum length of the input.
base (float, optional): The base for the frequency. Defaults to 10000.
Returns:
Tensor: The rotary embedding tensor.
"""
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64) / dim)
t = torch.arange(0, max_len, dtype=torch.float64)
freqs = torch.outer(t, theta)
return torch.cos(freqs).float(), torch.sin(freqs).float()
def apply_rotary_emb(x: torch.Tensor, rotary_emb: Tuple[Tensor, Tensor]) -> Tensor:
"""
Apply rotary embedding to the input tensor using cos/sin form.
Args:
x (Tensor): The input tensor (shape [..., seq_len, dim]).
rotary_emb (Tuple[Tensor, Tensor]): The rotary embedding (shape [seq_len, dim//2]).
Returns:
Tensor: The output tensor (rotated, same shape as input).
"""
dtype = x.dtype
cos, sin = rotary_emb
cos = cos.unsqueeze(0).unsqueeze(2) # [1, seq_len, 1, dim//2]
sin = sin.unsqueeze(0).unsqueeze(2) # [1, seq_len, 1, dim//2]
x_real = x[..., 0::2] # [batch, seq_len, dim//2]
x_imag = x[..., 1::2] # [batch, seq_len, dim//2]
x_real_rot = x_real * cos - x_imag * sin
x_imag_rot = x_real * sin + x_imag * cos
x_out = torch.stack([x_real_rot, x_imag_rot], dim=-1) # [batch, seq_len, dim//2, 2]
x_out = x_out.view(*x_out.shape[:-2], -1) # [batch, seq_len, dim]
return x_out.to(dtype)
class RotaryEmbedding(nn.Module):
def __init__(self, dim: int, max_len: int, base: int=10000):
super().__init__()
self.dim = dim
self.max_len = max_len
self.base = base
self.max_len_cached = None
self._set_rotary_buffer(self.max_len)
def _set_rotary_buffer(self, max_len: int):
cos_cached, sin_cached = get_rotary_emb(self.dim, max_len, self.base)
self.register_buffer("cos_cached", cos_cached, persistent=False)
self.register_buffer("sin_cached", sin_cached, persistent=False)
self.max_len_cached = max_len
def forward(self, x: Tensor, start_pos: int=0) -> Tuple[Tensor, Tensor]:
seq_len = x.size(1)
if self.max_len_cached < seq_len + start_pos:
self._set_rotary_buffer(seq_len)
cos = self.cos_cached[start_pos : start_pos + seq_len]
sin = self.sin_cached[start_pos : start_pos + seq_len]
return (cos, sin)
class Linear(nn.Module):
def __init__(self, in_dim: int, out_dim: int, bias: bool = False, weight_param=None, bias_param=None):
super().__init__()
weight_param = torch.empty((out_dim, in_dim)) if weight_param is None else weight_param
bias_param = torch.zeros(out_dim) if bias_param is None else bias_param
self.weight = nn.Parameter(weight_param)
self.bias = nn.Parameter(bias_param) if bias else None
def forward(self, x: Tensor) -> Tensor:
return F.linear(x, self.weight, self.bias)
class RMSNorm(nn.Module):
def __init__(self, n_dim, norm_eps):
super().__init__()
self.weight = nn.Parameter(torch.ones(n_dim))
self.norm_eps = norm_eps
def forward(self, x: Tensor) -> Tensor:
dtype = x.dtype
x = x.float()
mean_square = torch.mean(torch.pow(x, 2), dim=-1, keepdim=True)
norm = x * torch.rsqrt(mean_square + self.norm_eps)
norm = norm.to(dtype)
out = norm * self.weight
return out
class MLP(nn.Module):
def __init__(self, n_dim: int, d_ffn: int):
super().__init__()
self.up = Linear(n_dim, d_ffn)
self.gate = Linear(n_dim, d_ffn)
self.down = Linear(d_ffn, n_dim)
def forward(self, x: Tensor) -> Tensor:
gated = self.up(x) * F.silu(self.gate(x))
out = self.down(gated)
return out
class GQA(nn.Module):
def __init__(
self,
n_dim: int,
n_head: int,
n_kvhead: int,
layer_id: int
):
super().__init__()
assert n_dim % n_head == 0
assert n_head % n_kvhead == 0
self.head_dim = n_dim // n_head
self.layer_id = layer_id
self.n_dim = n_dim
self.n_heads = n_head
self.n_kvheads = n_kvhead
self.n_rep = n_head // n_kvhead
self.q_proj = Linear(n_dim, n_head * self.head_dim)
self.k_proj = Linear(n_dim, n_kvhead * self.head_dim)
self.v_proj = Linear(n_dim, n_kvhead * self.head_dim)
self.o_proj = Linear(n_dim, n_dim)
def forward(
self,
x: Tensor,
rotary_emb: Tuple[Tensor, Tensor],
mask: Tensor = None,
kv_cache: Optional[Tuple[Tensor, Tensor]] = None,
start_pos: int = 0
) -> Tensor:
bsz, seq_len, _ = x.size()
# x(bsz, seq_len, n_heads * head_dim) -> (bsz, seq_len, n_heads, head_dim)
q = self._split_heads(self.q_proj(x), self.n_heads)
k = self._split_heads(self.k_proj(x), self.n_kvheads)
v = self._split_heads(self.v_proj(x), self.n_kvheads)
q, k = apply_rotary_emb(q, rotary_emb), apply_rotary_emb(k, rotary_emb)
if kv_cache is not None:
k_cache, v_cache = kv_cache
# copy to cache
k_cache[:bsz, self.layer_id, start_pos:start_pos + seq_len] = k
v_cache[:bsz, self.layer_id, start_pos:start_pos + seq_len] = v
# get cache
k = k_cache[:bsz, self.layer_id, :start_pos + seq_len]
v = v_cache[:bsz, self.layer_id, :start_pos + seq_len]
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
# (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim)
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
sdqa_out = F.scaled_dot_product_attention(q, k, v, mask, is_causal=(mask == None)).permute(0, 2, 1, 3)
out = self.o_proj(sdqa_out.contiguous().view(bsz, seq_len, -1))
return out
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
batch_size, seq_len, _ = x.shape
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
return x
class DecoderBlock(nn.Module):
def __init__(self, n_dim, n_head, d_ffn, n_kvhead, norm_eps, layer_id):
super().__init__()
self.attention = GQA(n_dim, n_head, n_kvhead, layer_id)
self.norm_attn = RMSNorm(n_dim, norm_eps)
self.ffn = MLP(n_dim, d_ffn)
self.norm_ffn = RMSNorm(n_dim, norm_eps)
def forward(
self,
x: Tensor,
rotary_emb: Tuple[Tensor, Tensor],
attention_mask: Optional[Tensor] = None,
kv_cache: Optional[Tuple[Tensor, Tensor]] = None,
start_pos: int = 0
) -> Tensor:
# attention
attn_output = self.attention(
self.norm_attn(x),
rotary_emb,
attention_mask,
kv_cache,
start_pos
)
x = attn_output + x
# feed forward
x = self.ffn(self.norm_ffn(x)) + x
return x
class Embedding(nn.Module):
def __init__(self, vocab_size: int, embedding_dim: int, weight_param=None):
super().__init__()
weight_param = torch.empty((vocab_size, embedding_dim)) if weight_param is None else weight_param
self.weight = nn.Parameter(weight_param)
def forward(self, x: Tensor) -> Tensor:
return F.embedding(x, self.weight)
+132
View File
@@ -0,0 +1,132 @@
import torch
import torch.nn as nn
from torch import Tensor
from typing import Any, Mapping, Optional, Tuple
from khaosz.config.model_config import ModelConfig
from khaosz.model.module import Embedding, DecoderBlock, Linear, RMSNorm, RotaryEmbedding
def process_attention_mask(
seq_mask: Tensor,
input_tensor: Tensor,
start_pos: int = 0,
is_causal: bool = False,
) -> Tensor:
"""
Create attention mask for GQA
Args:
seq_mask (Tensor): A tensor indicating whether each position is valid or not.
input_tensor (Tensor): The input tensor.
start_pos (int): The starting position of the sequence.
is_causal (bool): Whether the attention is causal or not.
Returns:
Tensor: The attention mask tensor.
"""
device = input_tensor.device
dtype = input_tensor.dtype
seq_len = input_tensor.size(1)
if seq_mask is None:
if start_pos != 0:
# for single prompt chat
seq_mask = torch.ones((1, seq_len), dtype=torch.bool, device=device)
else:
return None
if seq_mask.dim() > 2:
# shape (bsz, seq_len) or (bsz,n_heads, seq_len, seq_len + start_pos)
# if ndim > 2, it's 4D tensor
return seq_mask
batch_size = seq_mask.size(0)
seq_mask = seq_mask[:, :start_pos + seq_len].to(device=device, dtype=torch.bool)
# (bsz, start_pos + seq_len)
expanded_mask = seq_mask.unsqueeze(1).expand(batch_size, seq_len, start_pos + seq_len)
# (bsz, seq_len, start_pos + seq_len)
if is_causal:
expanded_mask = torch.tril(expanded_mask, diagonal=start_pos)
attention_mask = torch.zeros_like(expanded_mask, dtype=dtype, device=device)
attention_mask = attention_mask.masked_fill_(~expanded_mask, -torch.finfo(dtype).max / 2).unsqueeze(1)
# (bsz, 1, seq_len, seq_len + start_pos)
return attention_mask
class Transformer(nn.Module):
def __init__(self, config: ModelConfig):
super().__init__()
self.config = config
self.rotary_embeding = RotaryEmbedding(config.n_dim // config.n_head, config.m_len)
self.embed_tokens = Embedding(config.vocab_size, config.n_dim)
self.layers = nn.ModuleList([
DecoderBlock(config.n_dim, config.n_head, config.d_ffn, config.n_kvhead, config.norm_eps, layer_id)
for layer_id in range(config.n_layer)
])
self.norm = RMSNorm(config.n_dim, config.norm_eps)
self.lm_head = Linear(config.n_dim, config.vocab_size)
if self.config.tie_weight == True:
self.lm_head.weight = self.embed_tokens.weight
self._init_parameters()
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
lm_head_key = 'lm_head.weight'
embed_key = 'embed_tokens.weight'
if self.config.tie_weight == True:
# same tensor
state_dict[lm_head_key] = state_dict[embed_key]
else:
# use clone to avoid sharing the same tensor
state_dict[lm_head_key] = torch.clone(state_dict[embed_key])
return super().load_state_dict(state_dict, strict, assign)
def state_dict(self, destination=None, prefix='', keep_vars=False):
state_dict = super().state_dict(destination=destination, prefix=prefix, keep_vars=keep_vars)
if self.config.tie_weight == True:
lm_head_key = prefix + 'lm_head.weight'
if lm_head_key in state_dict:
del state_dict[lm_head_key]
return state_dict
def _init_parameters(self):
for param in self.parameters():
if param.dim() > 1:
nn.init.normal_(param, mean=0.0, std=0.006)
def forward(
self,
input_ids: Tensor,
input_mask: Optional[Tensor]=None,
persistent_key_values: Optional[Tuple[Tensor, Tensor]]=None,
start_pos: int = 0
) -> Tensor:
assert input_ids.ndim == 2
x = self.embed_tokens(input_ids)
rotary_emb = self.rotary_embeding(x, start_pos)
attn_mask = process_attention_mask(
input_mask, x, start_pos, is_causal=True
)
for layer in self.layers:
x = layer(x, rotary_emb, attn_mask, persistent_key_values, start_pos)
hidden_states = self.norm(x)
logits = self.lm_head(hidden_states)
return {
"logits": logits,
"hidden_states": hidden_states
}
+14 -19
View File
@@ -1,34 +1,29 @@
from khaosz.trainer.data_util import DatasetLoader
from khaosz.trainer.trainer import Trainer from khaosz.trainer.trainer import Trainer
from khaosz.trainer.strategy import ( from khaosz.trainer.strategy import StrategyFactory
TrainConfig, from khaosz.trainer.schedule import SchedulerFactory
CosineScheduleConfig,
SgdrScheduleConfig, from khaosz.trainer.train_callback import (
StrategyFactory, TrainCallback,
SchedulerFactory
)
from khaosz.trainer.trainer_callback import (
TrainerCallback,
ProgressBarCallback, ProgressBarCallback,
CheckpointCallback, CheckpointCallback,
TrainerCallback, TrainCallback,
SchedulerCallback SchedulerCallback,
StepMonitorCallback
) )
__all__ = [ __all__ = [
# strategy # trainer
"DatasetLoader",
"Trainer", "Trainer",
"TrainConfig",
"CosineScheduleConfig", # factory
"SgdrScheduleConfig",
"StrategyFactory", "StrategyFactory",
"SchedulerFactory", "SchedulerFactory",
# callback # callback
"TrainerCallback", "TrainCallback",
"ProgressBarCallback", "ProgressBarCallback",
"CheckpointCallback", "CheckpointCallback",
"TrainerCallback", "TrainCallback",
"SchedulerCallback", "SchedulerCallback",
"StepMonitorCallback"
] ]
+65
View File
@@ -0,0 +1,65 @@
import torch.nn as nn
from typing import Dict
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
""" Compute gradient norm for each parameter in the model. """
norms = {}
for name, param in model.named_parameters():
norms[name] = 0.0
if param.grad:
norm = param.grad.data.norm(norm_type).item()
norms[name] = norm
return norms
def grad_std(model: nn.Module) -> Dict[str, float]:
""" Compute standard deviation of gradients for each parameter. """
stds = {}
for name, param in model.named_parameters():
stds[name] = 0.0
if param.grad:
std = param.grad.data.std().item()
stds[name] = std
return stds
def grad_max(model: nn.Module) -> Dict[str, float]:
""" Find the maximum absolute gradient value for each parameter. """
max_vals = {}
for name, param in model.named_parameters():
max_vals[name] = -float('inf')
if param.grad:
max_val = param.grad.data.max().item()
max_vals[name] = max_val
return max_vals
def grad_min(model: nn.Module) -> Dict[str, float]:
""" Find the minimum absolute gradient value for each parameter. """
min_vals = {}
for name, param in model.named_parameters():
min_vals[name] = float('inf')
if param.grad:
min_val = param.grad.data.min().item()
min_vals[name] = min_val
return min_vals
def grad_mean(model: nn.Module) -> Dict[str, float]:
""" Compute mean of gradients for each parameter. """
means = {}
for name, param in model.named_parameters():
means[name] = 0.0
if param.grad:
mean = param.grad.data.mean().item()
means[name] = mean
return means
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
""" Count the number of NaNs in gradients for each parameter. """
nan_nums = {}
for name, param in model.named_parameters():
nan_nums[name] = 0
if param.grad:
nan_num = param.grad.isnan().sum().item()
nan_nums[name] = nan_num
return nan_nums
+164
View File
@@ -0,0 +1,164 @@
import math
from abc import abstractmethod, ABC
from typing import Any, Dict, List
from torch.optim.lr_scheduler import LRScheduler
from khaosz.config.schedule_config import ScheduleConfig
class BaseScheduler(LRScheduler, ABC):
"""
Base scheduler class for all other schedulers.
"""
def __init__(self, optimizer, last_epoch: int = -1):
super().__init__(optimizer, last_epoch)
@abstractmethod
def get_lr(self) -> List[float]:
raise NotImplementedError
def state_dict(self) -> Dict[str, Any]:
return super().state_dict()
def load_state_dict(self, state_dict: Dict[str, Any]):
super().load_state_dict(state_dict)
class CosineScheduler(BaseScheduler):
"""
Cosine decay scheduler with warmup, implemented as PyTorch LRScheduler.
"""
def __init__(
self,
optimizer,
warmup_steps: int,
lr_decay_steps: int,
min_rate: float = 0.05,
last_epoch: int = -1
):
self.warmup_steps = warmup_steps
self.lr_decay_steps = lr_decay_steps
self.min_rate = min_rate
self.total_steps = warmup_steps + lr_decay_steps
super().__init__(optimizer, last_epoch)
def get_lr(self) -> List[float]:
# warmup
if self.last_epoch < self.warmup_steps:
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
return [base_lr * warmup_factor for base_lr in self.base_lrs]
# cosine decay
decay_progress = (self.last_epoch - self.warmup_steps) / self.lr_decay_steps
decay_progress = min(decay_progress, 1.0)
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * decay_progress))
decay_factor = max(self.min_rate, cosine_decay)
return [base_lr * decay_factor for base_lr in self.base_lrs]
def state_dict(self):
state = super().state_dict()
state.update({
'warmup_steps': self.warmup_steps,
'lr_decay_steps': self.lr_decay_steps,
'min_rate': self.min_rate,
'total_steps': self.total_steps,
})
return state
def load_state_dict(self, state_dict):
self.warmup_steps = state_dict.pop('warmup_steps')
self.lr_decay_steps = state_dict.pop('lr_decay_steps')
self.min_rate = state_dict.pop('min_rate')
self.total_steps = state_dict.pop('total_steps')
super().load_state_dict(state_dict)
class SGDRScheduler(BaseScheduler):
"""
SGDR (Stochastic Gradient Descent with Warm Restarts) scheduler,
"""
def __init__(
self,
optimizer,
warmup_steps: int,
cycle_length: int,
min_rate: float = 0.05,
t_mult: int = 2,
last_epoch: int = -1,
):
self.warmup_steps = warmup_steps
self.cycle_length = cycle_length
self.min_rate = min_rate
self.t_mult = t_mult
super().__init__(optimizer, last_epoch)
def get_lr(self):
# warmup
if self.last_epoch < self.warmup_steps:
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
return [base_lr * warmup_factor for base_lr in self.base_lrs]
# SGDR
steps_since_warmup = self.last_epoch - self.warmup_steps
# 1. Calculate current cycle and position within cycle
current_cycle_length = self.cycle_length
total_cycles_length = 0
cycle_num = 0
while total_cycles_length + current_cycle_length <= steps_since_warmup:
total_cycles_length += current_cycle_length
current_cycle_length *= self.t_mult
cycle_num += 1
steps_in_cycle = steps_since_warmup - total_cycles_length
# 2. Cosine annealing within the current cycle
cosine_factor = 0.5 * (1 + math.cos(math.pi * steps_in_cycle / current_cycle_length))
learning_rate_factor = self.min_rate + (1 - self.min_rate) * cosine_factor
return [base_lr * learning_rate_factor for base_lr in self.base_lrs]
def state_dict(self):
"""Returns the state of the scheduler as a dict."""
state = super().state_dict()
state.update({
'warmup_steps': self.warmup_steps,
'cycle_length': self.cycle_length,
'min_rate': self.min_rate,
't_mult': self.t_mult
})
return state
def load_state_dict(self, state_dict):
"""Loads the scheduler's state."""
self.warmup_steps = state_dict.pop('warmup_steps')
self.cycle_length = state_dict.pop('cycle_length')
self.min_rate = state_dict.pop('min_rate')
self.t_mult = state_dict.pop('t_mult')
super().load_state_dict(state_dict)
class SchedulerFactory:
"""
Factory class for creating learning rate schedulers.
"""
@staticmethod
def load_scheduler(optimizer, scedule_config: ScheduleConfig) -> BaseScheduler:
kwargs = scedule_config.get_kwargs()
schedule_type = kwargs.pop("schedule_type")
if schedule_type == "cosine":
return CosineScheduler(optimizer, **kwargs)
elif schedule_type == "sgdr":
return SGDRScheduler(optimizer, **kwargs)
else:
raise ValueError(f"Unsupported schedule type: {schedule_type}")
+26 -255
View File
@@ -1,15 +1,11 @@
import copy import copy
import math
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from torch.optim import Optimizer from typing import Any, Tuple, Callable, Dict, Union
from torch.utils.data import Dataset
from typing import Any, Literal, Tuple, Callable, Dict
from abc import ABC, abstractmethod 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: int): def get_logprobs(model:nn.Module, input_ids: Tensor, mask: Tensor, pad_token_id: int):
@@ -32,10 +28,14 @@ def get_logprobs(model:nn.Module, input_ids: Tensor, mask: Tensor, pad_token_id:
return (token_logprobs * valid_mask).sum(dim=-1) return (token_logprobs * valid_mask).sum(dim=-1)
def move_to_device(batch:Dict[str, Tensor], device: str) -> Any:
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
class BaseStrategy(ABC): class BaseStrategy(ABC):
def __init__(self, model: nn.Module): def __init__(self, model: Union[nn.Module, Callable[..., Dict[str, Tensor]]], device: str):
self.model = model self.model = model
self.device = device
@abstractmethod @abstractmethod
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
@@ -46,41 +46,36 @@ class BaseStrategy(ABC):
class SeqStrategy(BaseStrategy): class SeqStrategy(BaseStrategy):
def __init__(self, model): def __init__(self, model, device):
super().__init__(model) super().__init__(model, device)
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
batch = move_to_device(batch, self.device)
input_ids, target_ids = batch["input_ids"], batch["target_ids"] input_ids, target_ids = batch["input_ids"], batch["target_ids"]
B, L = input_ids.size() logits = self.model(input_ids=input_ids)["logits"]
logits: Tensor = self.model(input_ids=input_ids)["logits"]
loss = F.cross_entropy( loss = F.cross_entropy(
input=logits.view(B * L, -1), input=logits.flatten(0, 1),
target=target_ids.flatten() target=target_ids.flatten()
) )
return loss return loss
class SftStrategy(BaseStrategy): class SftStrategy(BaseStrategy):
def __init__(self, model: nn.Module): def __init__(self, model, device):
super().__init__(model) super().__init__(model, device)
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
input_ids, target_ids = batch["input_ids"], batch["target_ids"] batch = move_to_device(batch, self.device)
loss_mask, attn_mask = batch["loss_mask"], batch["attn_mask"] input_ids, target_ids, loss_mask = batch["input_ids"], batch["target_ids"], batch["loss_mask"]
ignore_index = -100 ignore_index = -100
B, L = input_ids.size() logits = self.model(input_ids=input_ids)["logits"]
logits: Tensor = self.model(
input_ids=input_ids,
input_mask=attn_mask
)["logits"]
target_ids = target_ids.masked_fill(loss_mask == 0, ignore_index) target_ids = target_ids.masked_fill(loss_mask == 0, ignore_index)
loss = F.cross_entropy( loss = F.cross_entropy(
input=logits.view(B * L, -1), input=logits.flatten(0, 1),
target=target_ids.flatten(), target=target_ids.flatten(),
ignore_index=ignore_index ignore_index=ignore_index
) )
@@ -89,8 +84,8 @@ class SftStrategy(BaseStrategy):
class DpoStrategy(BaseStrategy): class DpoStrategy(BaseStrategy):
def __init__(self, model, pad_token_id, beta): def __init__(self, model, device, pad_token_id, beta):
super().__init__(model) super().__init__(model, device)
ref_model = copy.deepcopy(self.model) ref_model = copy.deepcopy(self.model)
ref_model.requires_grad_(False) ref_model.requires_grad_(False)
ref_model.eval() ref_model.eval()
@@ -100,6 +95,7 @@ class DpoStrategy(BaseStrategy):
self.beta = beta self.beta = beta
def compute_loss(self, batch: Tuple[Tensor, ...]) -> Tensor: def compute_loss(self, batch: Tuple[Tensor, ...]) -> Tensor:
batch = move_to_device(batch, self.device)
good_ids, bad_ids = batch["chosen"], batch["rejected"] good_ids, bad_ids = batch["chosen"], batch["rejected"]
good_mask, bad_mask = batch["chosen_mask"], batch["rejected_mask"] good_mask, bad_mask = batch["chosen_mask"], batch["rejected_mask"]
@@ -156,241 +152,16 @@ class PpoStrategy(BaseStrategy):
class StrategyFactory: class StrategyFactory:
def load(model, train_type, **kwargs): def load(model, train_type, device, **kwargs):
train_strategy: Dict[str, Callable[[], BaseStrategy]] = { train_strategy: Dict[str, Callable[[], BaseStrategy]] = {
"seq": lambda: SeqStrategy(model), "seq": lambda: SeqStrategy(model, device),
"sft": lambda: SftStrategy(model), "sft": lambda: SftStrategy(model, device),
"dpo": lambda: DpoStrategy( "dpo": lambda: DpoStrategy(
model, model,
device,
kwargs.get("pad_token_id"), kwargs.get("pad_token_id"),
kwargs.get("dpo_beta") kwargs.get("dpo_beta")
) )
} }
strategy = train_strategy[train_type]() strategy = train_strategy[train_type]()
return strategy return strategy
@dataclass
class TrainConfig:
strategy: BaseStrategy = field(
default=None,
metadata={"help": "Training strategy."}
)
dataset: Dataset = field(
default=None,
metadata={"help": "Dataset for training."}
)
optimizer: Optimizer = field(
default=None,
metadata={"help": "Optimizer for training."}
)
checkpoint_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."}
)
checkpoint_interval: int = field(
default=5000,
metadata={"help": "Number of iterations between checkpoints."}
)
accumulation_steps: 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."}
)
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(ABC):
schedule_type: str = field(
default="cosine",
metadata={
"help": "Type of learning rate schedule.",
"choices": ["cosine", "sgdr"]
}
)
warmup_steps: int = field(
default=1000,
metadata={"help": "Number of warmup steps."}
)
min_rate: float = field(
default=0.05,
metadata={"help": "Minimum learning rate multiplier."}
)
@abstractmethod
def get_kwargs(self) -> Dict[str, Any]:
raise NotImplementedError
def validate(self) -> None:
"""Validate configuration parameters."""
if self.warmup_steps < 0:
raise ValueError(f"warmup_steps must be non-negative, got {self.warmup_steps}")
if not 0 <= self.min_rate <= 1:
raise ValueError(f"min_rate must be between 0 and 1, got {self.min_rate}")
@dataclass
class CosineScheduleConfig(ScheduleConfig):
total_steps: int = field(
default=None,
metadata={"help": "Total training steps for cosine schedule."}
)
schedule_type: Literal["cosine"] = "cosine"
def get_kwargs(self) -> Dict[str, Any]:
if self.total_steps is None:
raise ValueError("total_steps must be specified for cosine schedule")
return {
"schedule_type": self.schedule_type,
"warmup_steps": self.warmup_steps,
"lr_decay_steps": self.total_steps - self.warmup_steps,
"min_rate": self.min_rate
}
def validate(self) -> None:
super().validate()
if self.total_steps is not None and self.total_steps <= self.warmup_steps:
raise ValueError(f"total_steps ({self.total_steps}) must be greater than warmup_steps ({self.warmup_steps})")
@dataclass
class SgdrScheduleConfig(ScheduleConfig):
cycle_length: int = field(
default=1000,
metadata={"help": "Length of the first cycle in steps."}
)
t_mult: int = field(
default=2,
metadata={"help": "Multiplier for cycle length growth."}
)
schedule_type: Literal["sgdr"] = "sgdr"
def get_kwargs(self) -> Dict[str, Any]:
return {
"schedule_type": self.schedule_type,
"warmup_steps": self.warmup_steps,
"cycle_length": self.cycle_length,
"min_rate": self.min_rate,
"t_mult": self.t_mult
}
def validate(self) -> None:
super().validate()
if self.cycle_length <= 0:
raise ValueError(f"cycle_length must be positive, got {self.cycle_length}")
if self.t_mult < 1:
raise ValueError(f"t_mult must be >= 1, got {self.t_mult}")
class SchedulerFactory:
"""Factory for creating learning rate schedule functions."""
@staticmethod
def get_sgdr_schedule(
warmup_steps: int,
cycle_length: int,
min_rate: float = 0.05,
t_mult: int = 2
) -> Callable[[int], float]:
"""
Create SGDR (Stochastic Gradient Descent with Warm Restarts) schedule.
Args:
warmup_steps: Number of warmup steps
cycle_length: Length of the first cycle
min_rate: Minimum learning rate multiplier
t_mult: Cycle length multiplier
Returns:
Schedule function that takes current step and returns LR multiplier
"""
def sgdr_schedule(current_step: int) -> float:
# Warmup phase
if current_step < warmup_steps:
return max(min_rate, current_step / warmup_steps)
# SGDR phase
steps_since_warmup = current_step - warmup_steps
# Find current cycle and position within cycle
cycle_start = 0
current_cycle_length = cycle_length
cycle_index = 0
while steps_since_warmup >= cycle_start + current_cycle_length:
cycle_start += current_cycle_length
current_cycle_length *= t_mult
cycle_index += 1
position_in_cycle = steps_since_warmup - cycle_start
progress = position_in_cycle / current_cycle_length
# Cosine annealing within cycle
return max(min_rate, 0.5 * (1 + math.cos(math.pi * progress)))
return sgdr_schedule
@staticmethod
def get_cosine_schedule(
warmup_steps: int,
lr_decay_steps: int,
min_rate: float = 0.05
) -> Callable[[int], float]:
"""
Create cosine decay schedule with warmup.
Args:
warmup_steps: Number of warmup steps
lr_decay_steps: Number of steps for cosine decay after warmup
min_rate: Minimum learning rate multiplier
Returns:
Schedule function that takes current step and returns LR multiplier
"""
def cosine_schedule(current_step: int) -> float:
if current_step < warmup_steps:
# Linear warmup
return max(min_rate, current_step / warmup_steps)
else:
# Cosine decay
decay_progress = (current_step - warmup_steps) / lr_decay_steps
decay_progress = min(decay_progress, 1.0) # Clamp at 1.0
return max(min_rate, 0.5 * (1.0 + math.cos(math.pi * decay_progress)))
return cosine_schedule
@staticmethod
def load_schedule_fn(scedule_config: ScheduleConfig) -> Callable[[int], float]:
kwargs = scedule_config.get_kwargs()
schedule_type = kwargs.pop("schedule_type")
if schedule_type == "cosine":
return SchedulerFactory.get_cosine_schedule(**kwargs)
elif schedule_type == "sgdr":
return SchedulerFactory.get_sgdr_schedule(**kwargs)
else:
raise ValueError(f"Unsupported schedule type: {schedule_type}")
+226
View File
@@ -0,0 +1,226 @@
import os
import json
import time
from pathlib import Path
from tqdm import tqdm
from torch.nn.utils import clip_grad_norm_
from torch.optim.lr_scheduler import LambdaLR
from typing import List, Optional, Protocol, TYPE_CHECKING
from khaosz.config import ScheduleConfig
from khaosz.trainer.metric_util import (
grad_max,
grad_min,
grad_norm,
grad_mean,
grad_std,
grad_nan_num
)
if TYPE_CHECKING:
from khaosz.trainer.trainer import Trainer
from khaosz.trainer.train_context import TrainContext
class TrainCallback(Protocol):
"""
Callback interface for trainer.
"""
def on_train_begin(self, trainer: 'Trainer', context: 'TrainContext'):
""" Called at the beginning of training. """
def on_train_end(self, trainer: 'Trainer', context: 'TrainContext'):
""" Called at the end of training. """
def on_epoch_begin(self, trainer: 'Trainer', context: 'TrainContext'):
""" Called at the beginning of each epoch. """
def on_epoch_end(self, trainer: 'Trainer', context: 'TrainContext'):
""" Called at the end of each epoch. """
def on_step_begin(self, trainer: 'Trainer', context: 'TrainContext'):
""" Called at the beginning of each step. """
def on_step_end(self, trainer: 'Trainer', context: 'TrainContext'):
""" Called at the end of each step."""
def on_batch_begin(self, trainer: 'Trainer', context: 'TrainContext'):
""" Called at the beginning of each batch. """
def on_batch_end(self, trainer: 'Trainer', context: 'TrainContext'):
""" Called at the end of each batch. """
def on_error(self, trainer: 'Trainer', context: 'TrainContext'):
""" Called when an error occurs during training. """
class GradientClippingCallback(TrainCallback):
"""
Gradient clipping callback for trainer.
"""
def __init__(self, max_grad_norm: float):
self.max_grad_norm = max_grad_norm
def on_step_begin(self, trainer: 'Trainer', context: 'TrainContext'):
_ = context
clip_grad_norm_(trainer.parameter.model.parameters(), self.max_grad_norm)
class SchedulerCallback(TrainCallback):
"""
Scheduler callback for trainer.
"""
def __init__(self, schedule_config: ScheduleConfig):
self.schedule_config = schedule_config
self.scheduler: Optional[LambdaLR] = None
def on_train_begin(self, trainer: 'Trainer', context: 'TrainContext'):
for group in trainer.train_config.optimizer.param_groups:
if "initial_lr" not in group:
group["initial_lr"] = group["lr"]
self.scheduler = context.scheduler
def on_batch_end(self, trainer: 'Trainer', context: 'TrainContext'):
_ = trainer, context
if self.scheduler:
self.scheduler.step()
class CheckpointCallback(TrainCallback):
"""
Checkpoint callback for trainer.
"""
def __init__(self, checkpoint_interval: int):
self.checkpoint_interval = checkpoint_interval
self.last_ckpt_iter = 0
def _save_checkpoint(self, trainer: 'Trainer', context: 'TrainContext'):
save_path = os.path.join(trainer.train_config.checkpoint_dir, f"iter_{context.batch_iter}")
context.checkpoint.optimizer_state = context.optimizer.state_dict()
context.checkpoint.scheduler_state = context.scheduler.state_dict()
context.checkpoint.epoch = context.epoch
context.checkpoint.batch_iter = context.batch_iter
context.checkpoint.save(save_path)
self.last_ckpt_iter = context.batch_iter
def on_batch_end(self, trainer: 'Trainer', context: 'TrainContext'):
context.checkpoint.loss_list.append(context.loss)
if context.batch_iter - self.last_ckpt_iter >= self.checkpoint_interval:
self._save_checkpoint(trainer, context)
def on_train_end(self, trainer: 'Trainer', context: 'TrainContext'):
if context.batch_iter != self.last_ckpt_iter:
self._save_checkpoint(trainer, context)
class ProgressBarCallback(TrainCallback):
"""
Progress bar callback for trainer.
"""
def __init__(self):
self.progress_bar: tqdm = None
def on_epoch_begin(self, trainer: 'Trainer', context: 'TrainContext'):
self.progress_bar = tqdm(
context.dataloader,
desc=f"Epoch {context.epoch+1}/{trainer.train_config.n_epoch}",
dynamic_ncols=True
)
def on_batch_end(self, trainer: 'Trainer', context: 'TrainContext'):
_ = trainer
self.progress_bar.set_postfix({
"loss": f"{context.loss:.4f}",
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}"
})
self.progress_bar.update(1)
def on_epoch_end(self, trainer: 'Trainer', context: 'TrainContext'):
_ = trainer, context
if self.progress_bar:
self.progress_bar.close()
class StepMonitorCallback(TrainCallback):
"""
Customizable logger callback for trainer.
This callback provides flexible logging capabilities for training metrics,
supporting multiple log formats and custom log handlers.
"""
def __init__(
self,
log_dir: Optional[str] = None,
log_interval: int = 100,
metrics: Optional[List[str]] = None
):
"""
Args:
log_dir: Directory to save log files. If None, logs won't be saved to file.
log_interval: Log every N steps
metrics: List of metrics to log. Supported: ['loss', 'lr', 'grad_norm', 'grad_std',
grad_max', 'grad_min', 'grad_mean', 'grad_nan_num']
custom_handlers: List of custom log handler functions
json_log: Whether to save logs in JSON format
"""
self.log_dir = Path(log_dir) if log_dir else Path(os.getcwd()) / "logs"
self.log_interval = log_interval
self.metrics = metrics or ['loss', 'lr']
self.step_num = 0
self.log_dir.mkdir(parents=True, exist_ok=True)
def _handle_info(self, trainer: 'Trainer', context: 'TrainContext'):
""" Logs training information to console and file. """
log_data = {
"timestamp": time.strftime('%Y-%m-%d %H:%M:%S'),
"epoch": context.epoch,
"iter": context.batch_iter,
"metrics": self.metrics,
}
for metric in self.metrics:
if metric == 'loss':
log_data[metric] = context.loss
elif metric == 'lr':
log_data[metric] = context.optimizer.param_groups[-1]['lr']
elif metric == 'grad_norm':
log_data[metric] = grad_norm(trainer.parameter.model)
elif metric == 'grad_std':
log_data[metric] = grad_std(trainer.parameter.model)
elif metric == 'grad_max':
log_data[metric] = grad_max(trainer.parameter.model)
elif metric == 'grad_min':
log_data[metric] = grad_min(trainer.parameter.model)
elif metric == 'grad_mean':
log_data[metric] = grad_mean(trainer.parameter.model)
elif metric == 'grad_nan_num':
log_data[metric] = grad_nan_num(trainer.parameter.model)
else:
raise ValueError(f"Invalid metric: {metric}")
return log_data
def _handle_log(self, trainer: 'Trainer', context: 'TrainContext'):
""" Logs training information to console and file. """
log_data = self._handle_info(trainer, context)
try:
log_file = self.log_dir / f"log_epoch_{context.epoch}_iter_{context.batch_iter}.json"
with open(log_file, 'a') as f:
json.dump(log_data, f, indent=4)
except Exception:
raise
def on_step_end(self, trainer: 'Trainer', context: 'TrainContext'):
if self.step_num % self.log_interval == 0:
self._handle_log(trainer, context)
self.step_num += 1
+105
View File
@@ -0,0 +1,105 @@
from dataclasses import dataclass, field, fields
from typing import Optional, Self, TYPE_CHECKING
from torch.optim import Optimizer
from torch.utils.data import DataLoader
from khaosz.config import Checkpoint
from khaosz.data import ResumeableRandomSampler
from khaosz.trainer.schedule import BaseScheduler, SchedulerFactory
if TYPE_CHECKING:
from khaosz.trainer.trainer import Trainer
@dataclass
class TrainContext:
dataloader: DataLoader = field(default=None)
optimizer: Optimizer = field(default=None)
scheduler: BaseScheduler = field(default=None)
checkpoint: Checkpoint = field(default=None)
epoch: int = field(default=0)
batch_iter: int = field(default=0)
loss: float = field(default=0.0)
def asdict(self) -> dict:
return {field.name: getattr(self, field.name)
for field in fields(self)}
class TrainContextBuilder:
def __init__(self, trainer: 'Trainer'):
self.trainer = trainer
self._context: TrainContext = None
def with_checkpoint(self, checkpoint: Optional[Checkpoint]) -> Self:
self._context = TrainContext()
if checkpoint is None:
checkpoint = Checkpoint(
model=self.trainer.parameter.model,
tokenizer=self.trainer.parameter.tokenizer,
config=self.trainer.parameter.config,
)
else:
# resume from the assigned checkpoint or assigned iteration
self._context.epoch = max(checkpoint.epoch, self.trainer.train_config.start_epoch)
self._context.batch_iter = max(checkpoint.batch_iter, self.trainer.train_config.start_batch)
self._context.checkpoint = checkpoint
return self
def with_optimizer(self) -> Self:
if self._context is None:
raise RuntimeError("Must call with_checkpoint() before with_optimizer()")
optimizer = self.trainer.train_config.optimizer
if self._context.checkpoint and self._context.checkpoint.optimizer_state:
optimizer.load_state_dict(self._context.checkpoint.optimizer_state)
self._context.optimizer = optimizer
if self._context.checkpoint:
self._context.checkpoint.optimizer_state = optimizer.state_dict()
return self
def with_scheduler(self) -> Self:
if not hasattr(self._context, 'optimizer') or self._context.optimizer is None:
raise RuntimeError("Must call with_optimizer() before with_scheduler()")
optimizer = self.trainer.train_config.optimizer
schedule_config = self.trainer.schedule_config
scheduler = SchedulerFactory.load_scheduler(optimizer, schedule_config)
if self._context.checkpoint and self._context.checkpoint.scheduler_state:
scheduler.load_state_dict(self._context.checkpoint.scheduler_state)
self._context.scheduler = scheduler
if self._context.checkpoint:
self._context.checkpoint.scheduler_state = scheduler.state_dict()
return self
def with_dataloader(self) -> Self:
# fix: change batch level batch_iter to sample level offset
sampler_offset = self._context.batch_iter * self.trainer.train_config.batch_size
resumeable_sampler = ResumeableRandomSampler(
data_source=self.trainer.train_config.dataset,
start_epoch=self._context.epoch,
start_iter=sampler_offset,
seed=self.trainer.train_config.random_seed
)
dataloader = DataLoader(
self.trainer.train_config.dataset,
batch_size=self.trainer.train_config.batch_size,
sampler=resumeable_sampler,
num_workers=self.trainer.train_config.num_workers,
pin_memory=self.trainer.train_config.pin_memory,
prefetch_factor=self.trainer.train_config.prefetch_factor
)
self._context.dataloader = dataloader
return self
def build(self) -> TrainContext:
return self._context
+53 -98
View File
@@ -1,140 +1,95 @@
import logging import logging
from typing import Optional, List, cast from typing import Optional, List
from torch.utils.data import DataLoader from khaosz.config import (
ModelParameter,
from khaosz.core import ModelParameter, Checkpoint Checkpoint,
from khaosz.trainer.data_util import RandomSampler ScheduleConfig,
from khaosz.trainer.strategy import TrainConfig, ScheduleConfig TrainConfig
from khaosz.trainer.trainer_callback import ( )
TrainerCallback, from khaosz.trainer.train_callback import (
TrainCallback,
ProgressBarCallback, ProgressBarCallback,
CheckpointCallback, CheckpointCallback,
GradientClippingCallback, GradientClippingCallback,
SchedulerCallback SchedulerCallback
) )
from khaosz.trainer.train_context import TrainContext, TrainContextBuilder
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class Trainer: class Trainer:
def __init__( def __init__(
self, self,
parameter: ModelParameter, parameter: ModelParameter,
train_config: TrainConfig, train_config: TrainConfig,
schedule_config: ScheduleConfig, schedule_config: ScheduleConfig,
callbacks: Optional[List[TrainerCallback]] = None callbacks: Optional[List[TrainCallback]] = None
): ):
self.parameter = parameter self.parameter = parameter
self.train_config = train_config self.train_config = train_config
self.schedule_config = schedule_config self.schedule_config = schedule_config
self.callbacks = callbacks or self._get_default_callbacks() self.callbacks = callbacks or self._get_default_callbacks()
def _get_default_callbacks(self) -> List[TrainerCallback]: def _get_default_callbacks(self) -> List[TrainCallback]:
return [ return [
ProgressBarCallback(), ProgressBarCallback(),
CheckpointCallback(self.train_config.checkpoint_interval), CheckpointCallback(self.train_config.checkpoint_interval),
GradientClippingCallback(), GradientClippingCallback(self.train_config.max_grad_norm),
SchedulerCallback(self.schedule_config), SchedulerCallback(self.schedule_config),
] ]
def _set_train_kwargs(self, kwargs: dict):
seed = self.train_config.random_seed
sampler = RandomSampler(data_source=self.train_config.dataset, seed=seed)
optim = self.train_config.optimizer
checkpoint = cast(Checkpoint, kwargs.get('checkpoint', None))
if checkpoint is None: def _build_context(self, checkpoint: Optional[Checkpoint]) -> TrainContext:
checkpoint = Checkpoint( return (TrainContextBuilder(self)
model=self.parameter.model, .with_checkpoint(checkpoint)
tokenizer=self.parameter.tokenizer, .with_optimizer()
config=self.parameter.config, .with_scheduler()
sampler_state=None, .with_dataloader()
optim_state=None, .build())
loss_list=[]
) def _call_callbacks(self, method_name: str, context: TrainContext):
sampler_state = checkpoint.sampler_state
optim_state = checkpoint.optim_state
if sampler_state:
sampler.load_state_dict(sampler_state)
if optim_state:
optim.load_state_dict(optim_state)
checkpoint.optim_state = optim.state_dict()
checkpoint.sampler_state = sampler.state_dict()
dataloader = DataLoader(
self.train_config.dataset,
batch_size=self.train_config.batch_size,
sampler=sampler
)
kwargs["dataloader"] = dataloader
kwargs["optimizer"] = self.train_config.optimizer
kwargs["epoch"] = sampler.epoch
kwargs["current_iter"] = sampler.current_iter
kwargs["sampler"] = sampler
kwargs["checkpoint"] = checkpoint
def _call_callbacks(self, method_name: str, **kwargs):
for callback in self.callbacks: for callback in self.callbacks:
method = getattr(callback, method_name, None) method = getattr(callback, method_name, None)
if method: if method:
method(self, **kwargs) method(self, context)
def train( def train(self, checkpoint: Optional[Checkpoint] = None) -> Checkpoint:
self, context = self._build_context(checkpoint)
checkpoint: Optional[Checkpoint] = None self._call_callbacks('on_train_begin', context)
) -> Checkpoint:
# train
train_kwargs = {
'checkpoint': checkpoint,
'dataloader': None,
'optimizer': None,
'sampler': None,
'epoch': 0,
'current_iter': 0,
'loss': 0.0,
}
self._set_train_kwargs(train_kwargs)
self._call_callbacks('on_train_begin', **train_kwargs)
dataloader = train_kwargs['dataloader']
checkpoint = train_kwargs['checkpoint']
start_epoch = train_kwargs['epoch']
try: try:
self.parameter.model.train() self.parameter.model.train()
for epoch in range(start_epoch, self.train_config.n_epoch): # 1.epoch
# epoch for epoch in range(context.epoch, self.train_config.n_epoch):
train_kwargs["epoch"] = epoch context.epoch = epoch
self._call_callbacks('on_epoch_begin', **train_kwargs) self._call_callbacks('on_epoch_begin', context)
for batch in dataloader:
for batch in context.dataloader:
if train_kwargs["current_iter"] % self.train_config.accumulation_steps == 0: if context.batch_iter % self.train_config.accumulation_steps == 0:
# step # 2. step
self._call_callbacks('on_step_begin', **train_kwargs) self._call_callbacks('on_step_begin', context)
self.train_config.optimizer.step() self.train_config.optimizer.step()
self.train_config.optimizer.zero_grad() self.train_config.optimizer.zero_grad()
self._call_callbacks('on_step_end', **train_kwargs) self._call_callbacks('on_step_end', context)
# batch
self._call_callbacks('on_batch_begin', **train_kwargs)
loss = self.train_config.strategy(batch)
train_kwargs["loss"] = loss.item()
train_kwargs["current_iter"] += 1
loss.backward()
self._call_callbacks('on_batch_end', **train_kwargs) # 3. batch
self._call_callbacks('on_batch_begin', context)
loss = self.train_config.strategy(batch)
context.loss = loss.item()
context.batch_iter += 1
# to make the loss normalized by accumulation steps
normalized_loss = loss / self.train_config.accumulation_steps
normalized_loss.backward()
self._call_callbacks('on_batch_end', context)
self._call_callbacks('on_epoch_end', **train_kwargs) self._call_callbacks('on_epoch_end', context)
except Exception as e: except Exception as e:
logger.error(f"Training failed: {str(e)}", exc_info=True) logger.error(f"Training failed: {str(e)}", exc_info=True)
self._call_callbacks('on_error', context)
raise raise
finally: finally:
self._call_callbacks('on_train_end', **train_kwargs) self._call_callbacks('on_train_end', context)
return checkpoint return context.checkpoint
-169
View File
@@ -1,169 +0,0 @@
import os
import torch.optim as optim
from tqdm import tqdm
from torch.nn.utils import clip_grad_norm_
from torch.optim.lr_scheduler import LambdaLR
from typing import Optional, cast, TYPE_CHECKING
from khaosz.core.parameter import Checkpoint
from khaosz.trainer.data_util import RandomSampler
from khaosz.trainer.strategy import ScheduleConfig, SchedulerFactory
if TYPE_CHECKING:
from khaosz.trainer.trainer import Trainer
class TrainerCallback:
"""
Callback interface for trainer.
and we use '_' to ignore unused parameters.
"""
def on_train_begin(self, trainer: 'Trainer', **kwargs):
""" Called at the beginning of training. """
_ = trainer, kwargs
def on_train_end(self, trainer: 'Trainer', **kwargs):
""" Called at the end of training. """
_ = trainer, kwargs
def on_epoch_begin(self, trainer: 'Trainer', **kwargs):
""" Called at the beginning of each epoch. """
_ = trainer, kwargs
def on_epoch_end(self, trainer: 'Trainer', **kwargs):
""" Called at the end of each epoch. """
_ = trainer, kwargs
def on_batch_begin(self, trainer: 'Trainer', **kwargs):
""" Called at the beginning of each batch. """
_ = trainer, kwargs
def on_batch_end(self, trainer: 'Trainer', **kwargs):
""" Called at the end of each batch. """
_ = trainer, kwargs
def on_step_begin(self, trainer: 'Trainer', **kwargs):
""" Called at the beginning of each step. """
_ = trainer, kwargs
def on_step_end(self, trainer: 'Trainer', **kwargs):
""" Called at the end of each step."""
_ = trainer, kwargs
class ProgressBarCallback(TrainerCallback):
"""
Progress bar callback for trainer.
"""
def __init__(self):
self.progress_bar: tqdm = None
def on_epoch_begin(self, trainer: 'Trainer', **kwargs):
epoch = kwargs.get('epoch')
dataloader = kwargs.get('dataloader')
self.progress_bar = tqdm(
dataloader,
desc=f"Epoch {epoch+1}/{trainer.train_config.n_epoch}",
dynamic_ncols=True
)
def on_batch_end(self, trainer: 'Trainer', **kwargs):
_ = trainer
loss = kwargs.get('loss')
optimizer = cast(optim.Optimizer, kwargs.get('optimizer'))
self.progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"lr": f"{optimizer.param_groups[-1]['lr']:.2e}"
})
self.progress_bar.update(1)
def on_epoch_end(self, trainer: 'Trainer', **kwargs):
_ = trainer, kwargs
if self.progress_bar:
self.progress_bar.close()
class CheckpointCallback(TrainerCallback):
"""
Checkpoint callback for trainer.
"""
def __init__(self, checkpoint_interval: int):
self.checkpoint_interval = checkpoint_interval
self.last_ckpt_iter = 0
@staticmethod
def _save_checkpoint(trainer: 'Trainer', **kwargs):
current_iter = kwargs.get('current_iter')
random_sampler = cast(RandomSampler, kwargs.get('sampler'))
optimizer = cast(optim.Optimizer, kwargs.get('optimizer'))
checkpoint = cast(Checkpoint, kwargs.get('checkpoint'))
save_path = os.path.join(trainer.train_config.checkpoint_dir, f"iter_{current_iter}")
checkpoint.sampler_state = random_sampler.state_dict()
checkpoint.optim_state = optimizer.state_dict()
checkpoint.save(save_path)
def on_batch_end(self, trainer: 'Trainer', **kwargs):
current_iter = kwargs.get('current_iter')
checkpoint = cast(Checkpoint, kwargs.get('checkpoint'))
loss = kwargs.get('loss')
checkpoint.loss_list.append(loss)
if current_iter - self.last_ckpt_iter >= self.checkpoint_interval:
CheckpointCallback._save_checkpoint(trainer, **kwargs)
self.last_ckpt_iter = current_iter
def on_train_end(self, trainer: 'Trainer', **kwargs):
current_iter = kwargs.get('current_iter')
if current_iter != self.last_ckpt_iter:
CheckpointCallback._save_checkpoint(trainer, **kwargs)
self.last_ckpt_iter = current_iter
class GradientClippingCallback(TrainerCallback):
"""
Gradient clipping callback for trainer.
"""
def on_step_begin(self, trainer: 'Trainer', **kwargs):
_ = kwargs
clip_grad_norm_(
trainer.parameter.model.parameters(),
trainer.train_config.max_grad_norm
)
class SchedulerCallback(TrainerCallback):
"""
Scheduler callback for trainer.
"""
def __init__(self, schedule_config: ScheduleConfig):
self.schedule_config = schedule_config
self.scheduler: Optional[LambdaLR] = None
self.current_iter = 0
def on_train_begin(self, trainer: 'Trainer', **kwargs):
self.current_iter = kwargs.get('current_iter')
for group in trainer.train_config.optimizer.param_groups:
if "initial_lr" not in group:
group["initial_lr"] = group["lr"]
self.schedule_config.validate()
lambda_scheduler_fn = SchedulerFactory.load_schedule_fn(
self.schedule_config
)
self.scheduler = LambdaLR(
trainer.train_config.optimizer,
lambda_scheduler_fn,
last_epoch=self.current_iter - 1
)
def on_batch_end(self, trainer: 'Trainer', **kwargs):
_ = trainer, kwargs
if self.scheduler:
self.scheduler.step()
self.current_iter += 1
+2 -2
View File
@@ -14,9 +14,9 @@ def generate_text():
response = model.text_generate( response = model.text_generate(
query=query, query=query,
temperature=0.6, temperature=0.8,
top_p=0.95, top_p=0.95,
top_k=30 top_k=50
) )
print(response) print(response)
+2 -2
View File
@@ -13,9 +13,9 @@ def batch_generate():
responses = model.batch_generate( responses = model.batch_generate(
queries=inputs, queries=inputs,
temperature=0.7, temperature=0.8,
top_p=0.95, top_p=0.95,
top_k=30 top_k=50
) )
for q, r in zip(inputs, responses): for q, r in zip(inputs, responses):
+2 -2
View File
@@ -30,9 +30,9 @@ if __name__ == "__main__":
retrive_response = model.retrieve_generate( retrive_response = model.retrieve_generate(
retrieved=retrieved, retrieved=retrieved,
query=query, query=query,
temperature=0.7, temperature=0.8,
top_k=30,
top_p=0.95, top_p=0.95,
top_k=50
) )
print("retrive content:") print("retrive content:")
+5 -5
View File
@@ -10,19 +10,19 @@ def chat():
model_dir = os.path.join(PROJECT_ROOT, "params") model_dir = os.path.join(PROJECT_ROOT, "params")
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16) model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
histroy = [] history = []
while True: while True:
query = input(">> ") query = input(">> ")
if query == "!exit": if query == "!exit":
break break
response_size = 0 response_size = 0
for response, histroy in model.stream_generate( for response, history in model.stream_generate(
query=query, query=query,
history=histroy, history=history,
temperature=0.7, temperature=0.8,
top_p=0.95, top_p=0.95,
top_k=30 top_k=50
): ):
print(response[response_size:], end="", flush=True) print(response[response_size:], end="", flush=True)
response_size = len(response) response_size = len(response)
+3 -2
View File
@@ -1,4 +1,5 @@
import re import re
import khaosz
from setuptools import find_packages, setup from setuptools import find_packages, setup
@@ -8,11 +9,11 @@ with open("requirements.txt") as f:
setup( setup(
name="khaosz", name="khaosz",
version="1.2.0", version=khaosz.__version__,
packages=find_packages(), packages=find_packages(),
install_requires=required, install_requires=required,
dependency_links=[ dependency_links=[
"https://download.pytorch.org/whl/cu126", "https://download.pytorch.org/whl/cu126",
], ],
python_requires="==3.12.*", python_requires=">=3.12",
) )
+133
View File
@@ -0,0 +1,133 @@
import os
import json
import numpy as np
import tempfile
import shutil
import torch
import pytest
import matplotlib
from torch.utils.data import Dataset
from khaosz.config.model_config import ModelConfig
from khaosz.data.tokenizer import BpeTokenizer
from khaosz.model.transformer import Transformer
matplotlib.use("Agg")
class RandomDataset(Dataset):
def __init__(self, length=None, max_length=64, vocab_size=1000):
self.length = length or int(np.random.randint(100, 200))
self.max_length = max_length
self.vocab_size = vocab_size
def __len__(self):
return self.length
def __getitem__(self, idx):
return {
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,))
}
class MultiTurnDataset(Dataset):
def __init__(self, length=None, max_length=64, vocab_size=1000):
self.length = length or int(np.random.randint(100, 200))
self.max_length = max_length
self.vocab_size = vocab_size
def __len__(self):
return self.length
def __getitem__(self, idx):
input_ids = torch.randint(0, self.vocab_size, (self.max_length,))
target_ids = torch.randint(0, self.vocab_size, (self.max_length,))
loss_mask = torch.randint(0, 1, (self.max_length,))
return {
"input_ids": input_ids,
"target_ids": target_ids,
"loss_mask": loss_mask,
}
class EarlyStoppingDataset(Dataset):
def __init__(self, length=10, stop_after=5):
self.length = length
self.stop_after = stop_after
self.count = 0
def __len__(self):
return self.length
def __getitem__(self, idx):
self.count += 1
if self.count == self.stop_after:
raise RuntimeError("Simulated early stopping")
return {
"input_ids": torch.randint(0, 1000, (64,)),
"target_ids": torch.randint(0, 1000, (64,))
}
@pytest.fixture
def base_test_env(request: pytest.FixtureRequest):
func_name = request.function.__name__
test_dir = tempfile.mkdtemp(prefix=f"{func_name}_")
config_path = os.path.join(test_dir, "config.json")
n_dim_choices = [8, 16, 32]
n_head_choices = [2, 4]
n_dim = int(np.random.choice(n_dim_choices))
n_head = int(np.random.choice(n_head_choices))
n_kvhead = n_head // 2
d_ffn = n_dim * 2
config = {
"vocab_size": 1000,
"n_dim": n_dim,
"n_head": n_head,
"n_kvhead": n_kvhead,
"d_ffn": d_ffn,
"m_len": 1024,
"n_layer": 4,
"norm_eps": 1e-5
}
with open(config_path, 'w') as f:
json.dump(config, f)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
transformer_config = ModelConfig().load(config_path)
model = Transformer(transformer_config).to(device=device)
tokenizer = BpeTokenizer()
yield {
"device": device,
"test_dir": test_dir,
"config_path": config_path,
"transformer_config": transformer_config,
"model": model,
"tokenizer": tokenizer,
}
shutil.rmtree(test_dir)
@pytest.fixture
def random_dataset():
dataset = RandomDataset()
yield dataset
@pytest.fixture
def multi_turn_dataset():
dataset = MultiTurnDataset()
yield dataset
@pytest.fixture
def early_stopping_dataset():
dataset = EarlyStoppingDataset()
yield dataset
+58
View File
@@ -0,0 +1,58 @@
import torch
from khaosz.config import *
from khaosz.trainer import *
def test_callback_integration(base_test_env, random_dataset):
"""Test that all callbacks are properly integrated"""
optimizer = torch.optim.AdamW(base_test_env["model"].parameters())
train_config = TrainConfig(
dataset=random_dataset,
optimizer=optimizer,
checkpoint_dir=base_test_env["test_dir"],
n_epoch=1,
batch_size=2,
checkpoint_interval=3,
accumulation_steps=1,
max_grad_norm=1.0,
random_seed=42
)
schedule_config = CosineScheduleConfig(
warmup_steps=10,
total_steps=20
)
# Create custom callbacks to track calls
callback_calls = []
class TrackingCallback(TrainCallback):
def on_train_begin(self, trainer, context):
callback_calls.append('on_train_begin')
def on_batch_end(self, trainer, context):
callback_calls.append('on_batch_end')
def on_epoch_end(self, trainer, context):
callback_calls.append('on_epoch_end')
train_config.strategy = StrategyFactory.load(base_test_env["model"], "seq", base_test_env["device"])
model_parameter = ModelParameter(
base_test_env["model"],
base_test_env["tokenizer"],
base_test_env["transformer_config"]
)
trainer = Trainer(
model_parameter,
train_config,
schedule_config,
callbacks=[TrackingCallback(), ProgressBarCallback()]
)
trainer.train()
# Verify callbacks were called
assert 'on_train_begin' in callback_calls
assert 'on_batch_end' in callback_calls
assert 'on_epoch_end' in callback_calls
+67
View File
@@ -0,0 +1,67 @@
import os
import torch
import pickle
import numpy as np
from khaosz.trainer import *
from khaosz.data.data_util import *
def test_dataset_loader_random_paths(base_test_env):
"""Test dataset loader with multiple random paths"""
test_dir = base_test_env["test_dir"]
# Create multiple pkl files with random data
num_files = np.random.randint(2, 5)
pkl_paths = []
for i in range(num_files):
pkl_path = os.path.join(test_dir, f"test_data_{i}.pkl")
seq_length = np.random.randint(50, 100)
dummy_data = {
"sequence": torch.randint(0, 1000, (seq_length,)),
"chosen": torch.randint(0, 1000, (seq_length,)),
"rejected": torch.randint(0, 1000, (seq_length,)),
"chosen_mask": torch.ones(seq_length, dtype=torch.bool),
"rejected_mask": torch.ones(seq_length, dtype=torch.bool)
}
with open(pkl_path, "wb") as f:
pickle.dump(dummy_data, f)
pkl_paths.append(pkl_path)
# Test loading with multiple paths
loaded_dataset = DatasetLoader.load(
train_type="seq",
load_path=pkl_paths,
window_size=64,
)
assert loaded_dataset is not None
assert len(loaded_dataset) > 0
def test_dpo_strategy_with_random_data(base_test_env):
"""Test DPO strategy with randomized preference data"""
test_dir = base_test_env["test_dir"]
# Create DPO-style data
pkl_path = os.path.join(test_dir, "dpo_data.pkl")
seq_length = np.random.randint(40, 80)
dummy_data = {
"chosen": torch.randint(0, 1000, (seq_length,)),
"rejected": torch.randint(0, 1000, (seq_length,)),
"chosen_mask": torch.ones(seq_length, dtype=torch.bool),
"rejected_mask": torch.ones(seq_length, dtype=torch.bool)
}
with open(pkl_path, "wb") as f:
pickle.dump(dummy_data, f)
# Load DPO dataset
dpo_dataset = DatasetLoader.load(
train_type="dpo",
load_path=pkl_path,
window_size=64,
)
assert dpo_dataset is not None
assert hasattr(dpo_dataset, 'fetcher')
+41
View File
@@ -0,0 +1,41 @@
import torch
import numpy as np
from khaosz.config import *
from khaosz.trainer import *
def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
"""Simulate early stopping behavior"""
optimizer = torch.optim.AdamW(base_test_env["model"].parameters())
train_config = TrainConfig(
dataset=early_stopping_dataset,
optimizer=optimizer,
checkpoint_dir=base_test_env["test_dir"],
n_epoch=2,
batch_size=2,
checkpoint_interval=2,
accumulation_steps=2,
random_seed=np.random.randint(1e4),
)
train_config.strategy = StrategyFactory.load(base_test_env["model"], "seq", base_test_env["device"])
model_parameter = ModelParameter(
base_test_env["model"],
base_test_env["tokenizer"],
base_test_env["transformer_config"]
)
schedule_config = CosineScheduleConfig(warmup_steps=10, total_steps=20)
trainer = Trainer(model_parameter, train_config, schedule_config)
# Should handle early stopping gracefully
checkpoint = None
try:
checkpoint = trainer.train()
assert len(checkpoint.loss_list) == 2
except Exception:
# Handle any exceptions
pass
checkpoint = trainer.train(checkpoint)
assert len(checkpoint.loss_list) == 10
+21 -8
View File
@@ -5,13 +5,17 @@ import shutil
import pytest import pytest
import tempfile import tempfile
import safetensors.torch as st import safetensors.torch as st
from khaosz.core import * from khaosz.trainer import *
from khaosz.core.generator import EmbeddingEncoderCore, GeneratorCore from khaosz.config import *
from khaosz.model import *
from khaosz.data import *
from khaosz.inference.generator import EmbeddingEncoderCore, GeneratorCore
from tokenizers import pre_tokenizers from tokenizers import pre_tokenizers
@pytest.fixture @pytest.fixture
def test_env(): def test_env(request: pytest.FixtureRequest):
test_dir = tempfile.mkdtemp() func_name = request.function.__name__
test_dir = tempfile.mkdtemp(prefix=f"{func_name}_")
config_path = os.path.join(test_dir, "config.json") config_path = os.path.join(test_dir, "config.json")
tokenizer_path = os.path.join(test_dir, "tokenizer.json") tokenizer_path = os.path.join(test_dir, "tokenizer.json")
model_path = os.path.join(test_dir, "model.safetensors") model_path = os.path.join(test_dir, "model.safetensors")
@@ -34,7 +38,7 @@ def test_env():
tokenizer.train_from_iterator(sp_token_iter, config["vocab_size"], 1) tokenizer.train_from_iterator(sp_token_iter, config["vocab_size"], 1)
tokenizer.save(tokenizer_path) tokenizer.save(tokenizer_path)
transformer_config = TransformerConfig().load(config_path) transformer_config = ModelConfig().load(config_path)
model = Transformer(transformer_config) model = Transformer(transformer_config)
st.save_file(model.state_dict(), model_path) st.save_file(model.state_dict(), model_path)
@@ -97,7 +101,16 @@ def test_generator_core(test_env):
test_env["transformer_config"] test_env["transformer_config"]
) )
generator = GeneratorCore(parameter) generator = GeneratorCore(parameter)
logits, incr = generator.compute_logits(torch.randint(0, test_env["transformer_config"].vocab_size, (4, 10))) input_ids = torch.randint(0, test_env["transformer_config"].vocab_size, (4, 10))
next_token_id, cache_increase = generator.generate_iterator(
input_ids=input_ids,
temperature=0.8,
top_k=50,
top_p=0.95,
attn_mask=None,
kv_caches=None,
start_pos=0
)
assert logits.shape == (4, test_env["transformer_config"].vocab_size) assert next_token_id.shape == (4, 1)
assert incr == 10 assert cache_increase == 10
+52
View File
@@ -0,0 +1,52 @@
from khaosz.trainer import *
from khaosz.data.data_util import *
def test_random_sampler_consistency(random_dataset):
"""Test RandomSampler produces consistent results with same seed"""
dataset = random_dataset
# Create two samplers with same seed
sampler1 = ResumeableRandomSampler(dataset, seed=42)
sampler2 = ResumeableRandomSampler(dataset, seed=42)
indices1 = list(iter(sampler1))
indices2 = list(iter(sampler2))
assert indices1 == indices2
def test_random_sampler_different_seeds(random_dataset):
"""Test RandomSampler produces different results with different seeds"""
dataset = random_dataset
# Create two samplers with different seeds
sampler1 = ResumeableRandomSampler(dataset, seed=42)
sampler2 = ResumeableRandomSampler(dataset, seed=123)
indices1 = list(iter(sampler1))
indices2 = list(iter(sampler2))
# Very high probability they should be different
assert indices1 != indices2
def test_sampler_across_epochs(random_dataset):
"""Test sampler behavior across multiple epochs"""
dataset = random_dataset
n = len(dataset)
sampler = ResumeableRandomSampler(dataset, seed=42)
# Get indices for first epoch
epoch1_indices = list(iter(sampler))
assert len(epoch1_indices) == n
# Get indices for second epoch
epoch2_indices = list(iter(sampler))
assert len(epoch2_indices) == n
# Check that epochs have different order (should be random)
assert epoch1_indices != epoch2_indices
# Check that all indices are present in each epoch
assert set(epoch1_indices) == set(range(n))
assert set(epoch2_indices) == set(range(n))
+124
View File
@@ -0,0 +1,124 @@
import os
import json
import torch
import pytest
import tempfile
import safetensors.torch as st
from khaosz.model.transformer import Transformer
from khaosz.config.model_config import ModelConfig
@pytest.fixture
def transformer_test_env():
"""创建Transformer测试专用环境"""
test_dir = tempfile.mkdtemp(prefix="transformer_test_")
config_path = os.path.join(test_dir, "config.json")
config = {
"vocab_size": 1000,
"n_dim": 128,
"n_head": 4,
"n_kvhead": 2,
"d_ffn": 256,
"m_len": 64,
"n_layer": 2,
"norm_eps": 1e-5
}
with open(config_path, 'w') as f:
json.dump(config, f)
yield {
"test_dir": test_dir,
"config_path": config_path,
"config": config
}
if os.path.exists(test_dir):
try:
for file in os.listdir(test_dir):
os.remove(os.path.join(test_dir, file))
os.rmdir(test_dir)
except:
pass
def test_tie_weight_init(transformer_test_env):
config_path = transformer_test_env["config_path"]
config_data = transformer_test_env["config"].copy()
# case 1: tie weight
config_data["tie_weight"] = True
with open(config_path, 'w') as f:
json.dump(config_data, f)
config = ModelConfig().load(config_path)
model = Transformer(config)
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
assert model.lm_head.weight.data_ptr() == model.embed_tokens.weight.data_ptr()
original_weight = model.embed_tokens.weight.clone()
model.embed_tokens.weight.data[0, 0] = 100.0
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
assert not torch.equal(model.lm_head.weight, original_weight)
# case 2: not tie weight
config_data["tie_weight"] = False
with open(config_path, 'w') as f:
json.dump(config_data, f)
config = ModelConfig().load(config_path)
model = Transformer(config)
assert not torch.equal(model.lm_head.weight, model.embed_tokens.weight)
assert model.lm_head.weight.data_ptr() != model.embed_tokens.weight.data_ptr()
def test_model_save_load_with_tie_weight(transformer_test_env):
test_dir = transformer_test_env["test_dir"]
model_path = os.path.join(test_dir, "model.safetensors")
config_data = transformer_test_env["config"].copy()
# case 1: tie weight
config_data["tie_weight"] = True
config_path = os.path.join(test_dir, "config.json")
with open(config_path, 'w') as f:
json.dump(config_data, f)
config = ModelConfig().load(config_path)
original_model = Transformer(config)
st.save_file(original_model.state_dict(), model_path)
loaded_config = ModelConfig().load(config_path)
model = Transformer(loaded_config)
model.load_state_dict(st.load_file(model_path))
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
assert model.lm_head.weight.data_ptr() == model.embed_tokens.weight.data_ptr()
assert "lm_head.weight" not in model.state_dict()
# case 2: not tie weight
config_data["tie_weight"] = False
with open(config_path, 'w') as f:
json.dump(config_data, f)
config = ModelConfig().load(config_path)
original_model = Transformer(config)
st.save_file(original_model.state_dict(), model_path)
loaded_config = ModelConfig().load(config_path)
model = Transformer(loaded_config)
model.load_state_dict(st.load_file(model_path))
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
assert model.lm_head.weight.data_ptr() != model.embed_tokens.weight.data_ptr()
assert "lm_head.weight" in model.state_dict()
+86
View File
@@ -0,0 +1,86 @@
import torch
import numpy as np
from khaosz.config import *
from khaosz.trainer import *
from khaosz.data.data_util import *
def test_different_batch_sizes(base_test_env, random_dataset):
"""Test training with different batch sizes"""
batch_sizes = [1, 2, 4, 8]
for batch_size in batch_sizes:
optimizer = torch.optim.AdamW(base_test_env["model"].parameters())
train_config = TrainConfig(
dataset=random_dataset,
optimizer=optimizer,
checkpoint_dir=base_test_env["test_dir"],
n_epoch=1,
batch_size=batch_size,
checkpoint_interval=5,
accumulation_steps=1,
max_grad_norm=1.0,
random_seed=np.random.randint(1000)
)
assert train_config.batch_size == batch_size
def test_gradient_accumulation(base_test_env, random_dataset):
"""Test training with different gradient accumulation steps"""
accumulation_steps_list = [1, 2, 4]
for accumulation_steps in accumulation_steps_list:
optimizer = torch.optim.AdamW(base_test_env["model"].parameters())
train_config = TrainConfig(
dataset=random_dataset,
optimizer=optimizer,
checkpoint_dir=base_test_env["test_dir"],
n_epoch=1,
batch_size=2,
checkpoint_interval=10,
accumulation_steps=accumulation_steps,
max_grad_norm=1.0,
random_seed=42
)
schedule_config = CosineScheduleConfig(
warmup_steps=10,
total_steps=20
)
train_config.strategy = StrategyFactory.load(base_test_env["model"], "seq", base_test_env["device"])
model_parameter = ModelParameter(
base_test_env["model"],
base_test_env["tokenizer"],
base_test_env["transformer_config"]
)
trainer = Trainer(model_parameter, train_config, schedule_config)
trainer.train()
assert train_config.accumulation_steps == accumulation_steps
def test_memory_efficient_training(base_test_env, random_dataset):
"""Test training with memory-efficient configurations"""
# Test with smaller batch sizes and gradient checkpointing
small_batch_configs = [
{"batch_size": 1, "accumulation_steps": 8},
{"batch_size": 2, "accumulation_steps": 4},
{"batch_size": 4, "accumulation_steps": 2}
]
for config in small_batch_configs:
optimizer = torch.optim.AdamW(base_test_env["model"].parameters())
train_config = TrainConfig(
dataset=random_dataset,
optimizer=optimizer,
checkpoint_dir=base_test_env["test_dir"],
n_epoch=1,
batch_size=config["batch_size"],
checkpoint_interval=5,
accumulation_steps=config["accumulation_steps"],
max_grad_norm=1.0,
random_seed=42
)
assert train_config.accumulation_steps == config["accumulation_steps"]
+135
View File
@@ -0,0 +1,135 @@
import torch
import numpy as np
import pytest
from khaosz.config import *
from khaosz.trainer.schedule import *
from khaosz.data.data_util import *
def test_schedule_factory_random_configs():
"""Test scheduler factory with random configurations"""
# Create a simple model and optimizer for testing
model = torch.nn.Linear(10, 2)
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001)
# Test multiple random configurations
for _ in range(5): # Test 5 random configurations
schedule_configs = [
CosineScheduleConfig(
warmup_steps=np.random.randint(50, 200),
total_steps=np.random.randint(1000, 5000),
min_rate=np.random.uniform(0.01, 0.1)
),
SGDRScheduleConfig(
warmup_steps=np.random.randint(50, 200),
cycle_length=np.random.randint(500, 2000),
t_mult=np.random.randint(1, 3),
min_rate=np.random.uniform(0.01, 0.1)
)
]
for config in schedule_configs:
# Validate configuration
config.validate()
# Create scheduler using factory
scheduler = SchedulerFactory.load_scheduler(optimizer, config)
# Verify scheduler type
if isinstance(config, CosineScheduleConfig):
assert isinstance(scheduler, CosineScheduler)
assert scheduler.warmup_steps == config.warmup_steps
assert scheduler.lr_decay_steps == config.total_steps - config.warmup_steps
assert scheduler.min_rate == config.min_rate
elif isinstance(config, SGDRScheduleConfig):
assert isinstance(scheduler, SGDRScheduler)
assert scheduler.warmup_steps == config.warmup_steps
assert scheduler.cycle_length == config.cycle_length
assert scheduler.t_mult == config.t_mult
assert scheduler.min_rate == config.min_rate
# Test scheduler state dict functionality
state_dict = scheduler.state_dict()
assert 'warmup_steps' in state_dict
assert 'min_rate' in state_dict
# Test scheduler step functionality
initial_lr = scheduler.get_last_lr()
scheduler.step()
new_lr = scheduler.get_last_lr()
# Learning rate should change after step, or if it's the first step,
# the epoch counter should increment
assert initial_lr != new_lr or scheduler.last_epoch > -1
def test_schedule_factory_edge_cases():
"""Test scheduler factory with edge cases and boundary conditions"""
model = torch.nn.Linear(10, 2)
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001)
# Test edge cases for CosineScheduleConfig
edge_cases = [
# Minimal warmup and steps
CosineScheduleConfig(warmup_steps=1, total_steps=10, min_rate=0.01),
# Large values
CosineScheduleConfig(warmup_steps=1000, total_steps=10000, min_rate=0.5),
# Zero min_rate (edge case)
CosineScheduleConfig(warmup_steps=100, total_steps=1000, min_rate=0.0),
]
for config in edge_cases:
config.validate()
scheduler = SchedulerFactory.load_scheduler(optimizer, config)
assert scheduler is not None
# Test multiple steps
for _ in range(10):
scheduler.step()
def test_schedule_factory_invalid_configs():
"""Test scheduler factory with invalid configurations"""
# Test invalid configurations that should raise errors
invalid_configs = [
# Negative warmup steps
CosineScheduleConfig(warmup_steps=-10, total_steps=1000, min_rate=0.1),
# Total steps less than warmup steps
CosineScheduleConfig(warmup_steps=500, total_steps=400, min_rate=0.1),
# Invalid min_rate
CosineScheduleConfig(warmup_steps=100, total_steps=1000, min_rate=-0.1),
CosineScheduleConfig(warmup_steps=100, total_steps=1000, min_rate=1.1),
]
for config in invalid_configs:
with pytest.raises(ValueError):
config.validate()
def test_schedule_factory_state_persistence():
"""Test scheduler state persistence (save/load)"""
model = torch.nn.Linear(10, 2)
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001)
config = CosineScheduleConfig(warmup_steps=100, total_steps=1000, min_rate=0.1)
scheduler = SchedulerFactory.load_scheduler(optimizer, config)
# Take a few steps
for _ in range(5):
scheduler.step()
# Save state
state_dict = scheduler.state_dict()
# Create new scheduler and load state
new_scheduler = SchedulerFactory.load_scheduler(optimizer, config)
new_scheduler.load_state_dict(state_dict)
# Verify states match
assert scheduler.last_epoch == new_scheduler.last_epoch
assert scheduler.get_last_lr() == new_scheduler.get_last_lr()
-466
View File
@@ -1,466 +0,0 @@
import os
import json
import torch
import shutil
import pytest
import pickle
import tempfile
import numpy as np
from torch.utils.data import Dataset
from khaosz.core import *
from khaosz.trainer import *
from khaosz.trainer.data_util import *
import matplotlib
matplotlib.use('Agg')
@pytest.fixture
def test_env():
"""Setup test environment with randomized data"""
test_dir = tempfile.mkdtemp()
config_path = os.path.join(test_dir, "config.json")
n_dim_choices = [8, 16, 32]
n_head_choices = [2, 4]
n_dim = int(np.random.choice(n_dim_choices))
n_head = int(np.random.choice(n_head_choices))
n_kvhead = n_head // 2
d_ffn = n_dim * 2
config = {
"vocab_size": 1000,
"n_dim": n_dim,
"n_head": n_head,
"n_kvhead": n_kvhead,
"d_ffn": d_ffn,
"m_len": 1024,
"n_layer": 4,
"norm_eps": 1e-5
}
with open(config_path, 'w') as f:
json.dump(config, f)
transformer_config = TransformerConfig().load(config_path)
model = Transformer(transformer_config)
tokenizer = BpeTokenizer()
class RandomDataset(Dataset):
def __init__(self, length=None, max_length=64, vocab_size=1000):
self.length = length or int(np.random.randint(100, 200))
self.max_length = max_length
self.vocab_size = vocab_size
def __len__(self):
return self.length
def __getitem__(self, idx):
return {
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,))
}
class MultiTurnDataset(Dataset):
def __init__(self, length=None, max_length=64, vocab_size=1000):
self.length = length or int(np.random.randint(100, 200))
self.max_length = max_length
self.vocab_size = vocab_size
def __len__(self):
return self.length
def __getitem__(self, idx):
input_ids = torch.randint(0, self.vocab_size, (self.max_length,))
target_ids = torch.randint(0, self.vocab_size, (self.max_length,))
loss_mask = build_loss_mask(input_ids, 0, 1)
attn_mask = build_attention_mask(input_ids, 2, True)
return {
"input_ids": input_ids,
"target_ids": target_ids,
"loss_mask": loss_mask,
"attn_mask": attn_mask,
}
dataset = RandomDataset()
multi_turn_dataset = MultiTurnDataset()
yield {
"test_dir": test_dir,
"config_path": config_path,
"transformer_config": transformer_config,
"model": model,
"tokenizer": tokenizer,
"dataset": dataset,
"multi_turn_dataset": multi_turn_dataset
}
shutil.rmtree(test_dir)
def test_dataset_loader_random_paths(test_env):
"""Test dataset loader with multiple random paths"""
test_dir = test_env["test_dir"]
# Create multiple pkl files with random data
num_files = np.random.randint(2, 5)
pkl_paths = []
for i in range(num_files):
pkl_path = os.path.join(test_dir, f"test_data_{i}.pkl")
seq_length = np.random.randint(50, 100)
dummy_data = {
"sequence": torch.randint(0, 1000, (seq_length,)),
"chosen": torch.randint(0, 1000, (seq_length,)),
"rejected": torch.randint(0, 1000, (seq_length,)),
"chosen_mask": torch.ones(seq_length, dtype=torch.bool),
"rejected_mask": torch.ones(seq_length, dtype=torch.bool)
}
with open(pkl_path, "wb") as f:
pickle.dump(dummy_data, f)
pkl_paths.append(pkl_path)
# Test loading with multiple paths
loaded_dataset = DatasetLoader.load(
train_type="seq",
load_path=pkl_paths,
max_len=64,
device="cpu"
)
assert loaded_dataset is not None
assert len(loaded_dataset) > 0
def test_different_batch_sizes(test_env):
"""Test training with different batch sizes"""
batch_sizes = [1, 2, 4, 8]
for batch_size in batch_sizes:
optimizer = torch.optim.AdamW(test_env["model"].parameters())
train_config = TrainConfig(
dataset=test_env["dataset"],
optimizer=optimizer,
checkpoint_dir=test_env["test_dir"],
n_epoch=1,
batch_size=batch_size,
checkpoint_interval=5,
accumulation_steps=1,
max_grad_norm=1.0,
random_seed=np.random.randint(1000)
)
assert train_config.batch_size == batch_size
def test_random_sampler_consistency(test_env):
"""Test RandomSampler produces consistent results with same seed"""
dataset = test_env["dataset"]
# Create two samplers with same seed
sampler1 = RandomSampler(dataset, seed=42)
sampler2 = RandomSampler(dataset, seed=42)
indices1 = list(iter(sampler1))
indices2 = list(iter(sampler2))
assert indices1 == indices2
def test_random_sampler_different_seeds(test_env):
"""Test RandomSampler produces different results with different seeds"""
dataset = test_env["dataset"]
# Create two samplers with different seeds
sampler1 = RandomSampler(dataset, seed=42)
sampler2 = RandomSampler(dataset, seed=123)
indices1 = list(iter(sampler1))
indices2 = list(iter(sampler2))
# Very high probability they should be different
assert indices1 != indices2
def test_schedule_factory_random_configs(test_env):
"""Test scheduler factory with random configurations"""
schedule_configs = [
CosineScheduleConfig(
warmup_steps=np.random.randint(50, 200),
total_steps=np.random.randint(1000, 5000),
min_rate=np.random.uniform(0.01, 0.1)
),
SgdrScheduleConfig(
warmup_steps=np.random.randint(50, 200),
cycle_length=np.random.randint(500, 2000),
t_mult=np.random.randint(1, 3),
min_rate=np.random.uniform(0.01, 0.1)
)
]
for config in schedule_configs:
schedule_fn = SchedulerFactory.load_schedule_fn(config)
assert callable(schedule_fn)
# Test the schedule function at different steps
for step in [0, config.warmup_steps // 2, config.warmup_steps, config.warmup_steps * 2]:
lr_mult = schedule_fn(step)
assert 0 <= lr_mult <= 1
def test_multi_turn_training(test_env):
"""Test training with multi-turn conversation data"""
optimizer = torch.optim.AdamW(test_env["model"].parameters())
train_config = TrainConfig(
dataset=test_env["multi_turn_dataset"],
optimizer=optimizer,
checkpoint_dir=test_env["test_dir"],
n_epoch=2,
batch_size=2,
checkpoint_interval=3,
accumulation_steps=1,
max_grad_norm=1.0,
random_seed=int(np.random.randint(1000))
)
schedule_config = CosineScheduleConfig(
warmup_steps=50,
total_steps=100
)
train_config.strategy = StrategyFactory.load(
test_env["model"],
"sft",
bos_token_id=2,
eos_token_id=3,
user_token_id=1,
multi_turn=True
)
model_parameter = ModelParameter(
test_env["model"],
test_env["tokenizer"],
test_env["transformer_config"]
)
trainer = Trainer(model_parameter, train_config, schedule_config)
checkpoint = trainer.train()
assert len(checkpoint.loss_list) > 0
def test_gradient_accumulation(test_env):
"""Test training with different gradient accumulation steps"""
accumulation_steps_list = [1, 2, 4]
for accumulation_steps in accumulation_steps_list:
optimizer = torch.optim.AdamW(test_env["model"].parameters())
train_config = TrainConfig(
dataset=test_env["dataset"],
optimizer=optimizer,
checkpoint_dir=test_env["test_dir"],
n_epoch=1,
batch_size=2,
checkpoint_interval=10,
accumulation_steps=accumulation_steps,
max_grad_norm=1.0,
random_seed=42
)
schedule_config = CosineScheduleConfig(
warmup_steps=10,
total_steps=20
)
train_config.strategy = StrategyFactory.load(
test_env["model"],
"seq"
)
model_parameter = ModelParameter(
test_env["model"],
test_env["tokenizer"],
test_env["transformer_config"]
)
trainer = Trainer(model_parameter, train_config, schedule_config)
trainer.train()
assert train_config.accumulation_steps == accumulation_steps
def test_dpo_strategy_with_random_data(test_env):
"""Test DPO strategy with randomized preference data"""
test_dir = test_env["test_dir"]
# Create DPO-style data
pkl_path = os.path.join(test_dir, "dpo_data.pkl")
seq_length = np.random.randint(40, 80)
dummy_data = {
"chosen": torch.randint(0, 1000, (seq_length,)),
"rejected": torch.randint(0, 1000, (seq_length,)),
"chosen_mask": torch.ones(seq_length, dtype=torch.bool),
"rejected_mask": torch.ones(seq_length, dtype=torch.bool)
}
with open(pkl_path, "wb") as f:
pickle.dump(dummy_data, f)
# Load DPO dataset
dpo_dataset = DatasetLoader.load(
train_type="dpo",
load_path=pkl_path,
max_len=64,
device="cpu"
)
assert dpo_dataset is not None
assert hasattr(dpo_dataset, 'fetcher')
def test_callback_integration(test_env):
"""Test that all callbacks are properly integrated"""
optimizer = torch.optim.AdamW(test_env["model"].parameters())
train_config = TrainConfig(
dataset=test_env["dataset"],
optimizer=optimizer,
checkpoint_dir=test_env["test_dir"],
n_epoch=1,
batch_size=2,
checkpoint_interval=3,
accumulation_steps=1,
max_grad_norm=1.0,
random_seed=42
)
schedule_config = CosineScheduleConfig(
warmup_steps=10,
total_steps=20
)
# Create custom callbacks to track calls
callback_calls = []
class TrackingCallback(TrainerCallback):
def on_train_begin(self, trainer, **kwargs):
callback_calls.append('on_train_begin')
def on_batch_end(self, trainer, **kwargs):
callback_calls.append('on_batch_end')
def on_epoch_end(self, trainer, **kwargs):
callback_calls.append('on_epoch_end')
train_config.strategy = StrategyFactory.load(test_env["model"], "seq")
model_parameter = ModelParameter(
test_env["model"],
test_env["tokenizer"],
test_env["transformer_config"]
)
trainer = Trainer(
model_parameter,
train_config,
schedule_config,
callbacks=[TrackingCallback(), ProgressBarCallback()]
)
trainer.train()
# Verify callbacks were called
assert 'on_train_begin' in callback_calls
assert 'on_batch_end' in callback_calls
assert 'on_epoch_end' in callback_calls
def test_memory_efficient_training(test_env):
"""Test training with memory-efficient configurations"""
# Test with smaller batch sizes and gradient checkpointing
small_batch_configs = [
{"batch_size": 1, "accumulation_steps": 8},
{"batch_size": 2, "accumulation_steps": 4},
{"batch_size": 4, "accumulation_steps": 2}
]
for config in small_batch_configs:
optimizer = torch.optim.AdamW(test_env["model"].parameters())
train_config = TrainConfig(
dataset=test_env["dataset"],
optimizer=optimizer,
checkpoint_dir=test_env["test_dir"],
n_epoch=1,
batch_size=config["batch_size"],
checkpoint_interval=5,
accumulation_steps=config["accumulation_steps"],
max_grad_norm=1.0,
random_seed=42
)
assert train_config.accumulation_steps == config["accumulation_steps"]
def test_early_stopping_simulation(test_env):
"""Simulate early stopping behavior"""
class EarlyStoppingDataset(Dataset):
def __init__(self, length=10, stop_after=5):
self.length = length
self.stop_after = stop_after
self.count = 0
def __len__(self):
return self.length
def __getitem__(self, idx):
self.count += 1
if self.count == self.stop_after:
raise RuntimeError("Simulated early stopping")
return {
"input_ids": torch.randint(0, 1000, (64,)),
"target_ids": torch.randint(0, 1000, (64,))
}
dataset = EarlyStoppingDataset()
optimizer = torch.optim.AdamW(test_env["model"].parameters())
train_config = TrainConfig(
dataset=dataset,
optimizer=optimizer,
checkpoint_dir=test_env["test_dir"],
n_epoch=2,
batch_size=2,
checkpoint_interval=1,
accumulation_steps=1,
max_grad_norm=1.0,
random_seed=42
)
train_config.strategy = StrategyFactory.load(test_env["model"], "seq")
model_parameter = ModelParameter(
test_env["model"],
test_env["tokenizer"],
test_env["transformer_config"]
)
schedule_config = CosineScheduleConfig(warmup_steps=10, total_steps=20)
trainer = Trainer(model_parameter, train_config, schedule_config)
# Should handle early stopping gracefully
checkpoint = None
try:
checkpoint = trainer.train()
assert len(checkpoint.loss_list) == 2
except Exception:
# Handle any exceptions
pass
checkpoint = trainer.train(checkpoint)
assert len(checkpoint.loss_list) == 10 + 1
if __name__ == "__main__":
# Run all tests
pytest.main([__file__, "-v"])
+30 -31
View File
@@ -3,9 +3,9 @@ import argparse
import torch import torch
from torch.optim import AdamW from torch.optim import AdamW
from khaosz.core import ParameterLoader from khaosz.config import ParameterLoader, Checkpoint, TrainConfig, CosineScheduleConfig
from khaosz.trainer import Trainer, DatasetLoader, TrainConfig, CosineScheduleConfig from khaosz.trainer import Trainer, StrategyFactory
from khaosz.trainer import StrategyFactory from khaosz.data import DatasetLoader
PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__)) PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
@@ -24,6 +24,8 @@ def train(
max_lr: int, max_lr: int,
n_epoch: int, n_epoch: int,
batch_size: int, batch_size: int,
start_epoch: int,
start_batch: int,
accumulation_steps: int, accumulation_steps: int,
warmup_steps: int, warmup_steps: int,
checkpoint_interval: int, checkpoint_interval: int,
@@ -34,39 +36,46 @@ def train(
max_grad_norm: float, max_grad_norm: float,
embdeding_lr_rate: int, embdeding_lr_rate: int,
random_seed: int, random_seed: int,
multi_turn: bool, window_size: int,
stride: int,
resume_from_checkpoint: bool
): ):
assert train_type in ["seq", "sft", "dpo"] assert train_type in ["seq", "sft", "dpo"]
assert os.path.exists(param_path) assert os.path.exists(param_path)
parameter = ParameterLoader.load(param_path) parameter = ParameterLoader.load(param_path)
model = parameter.model checkpoint = None
if isinstance(parameter, Checkpoint) and resume_from_checkpoint:
checkpoint = parameter
if window_size is None:
window_size = parameter.config.m_len
model = parameter.model
device = torch.device("cuda") device = torch.device("cuda")
model = model.to(device=device, dtype=torch.bfloat16) model = model.to(device=device, dtype=torch.bfloat16)
cache_files = get_files(data_root_path) cache_files = get_files(data_root_path)
kwargs = { kwargs = {
"multi_turn": multi_turn,
"dpo_beta": dpo_beta, "dpo_beta": dpo_beta,
"bos_token_id": parameter.tokenizer.bos_id, "bos_token_id": parameter.tokenizer.bos_id,
"eos_token_id": parameter.tokenizer.eos_id, "eos_token_id": parameter.tokenizer.eos_id,
"pad_token_id": parameter.tokenizer.pad_id, "pad_token_id": parameter.tokenizer.pad_id,
"user_token_id":parameter.tokenizer.user_id,
} }
strategy = StrategyFactory.load( strategy = StrategyFactory.load(
model, model,
train_type, train_type,
device,
**kwargs **kwargs
) )
dataset = DatasetLoader.load( dataset = DatasetLoader.load(
train_type=train_type, train_type=train_type,
load_path=cache_files, load_path=cache_files,
max_len=parameter.config.m_len, window_size=window_size,
device=device, stride=stride,
**kwargs **kwargs
) )
@@ -88,10 +97,14 @@ def train(
checkpoint_dir=checkpoint_dir, checkpoint_dir=checkpoint_dir,
n_epoch=n_epoch, n_epoch=n_epoch,
batch_size=batch_size, batch_size=batch_size,
start_epoch=start_epoch,
start_batch=start_batch,
checkpoint_interval=checkpoint_interval, checkpoint_interval=checkpoint_interval,
accumulation_steps=accumulation_steps, accumulation_steps=accumulation_steps,
max_grad_norm=max_grad_norm, max_grad_norm=max_grad_norm,
random_seed=random_seed, random_seed=random_seed,
num_workers=4,
pin_memory=True
) )
schedule_config = CosineScheduleConfig( schedule_config = CosineScheduleConfig(
@@ -104,7 +117,7 @@ def train(
train_config=train_config, train_config=train_config,
schedule_config=schedule_config, schedule_config=schedule_config,
) )
trainer.train() trainer.train(checkpoint)
if __name__ == "__main__": if __name__ == "__main__":
@@ -127,27 +140,13 @@ if __name__ == "__main__":
parser.add_argument("--random_seed", type=int, default=3407, help="Random seed for reproducibility.") parser.add_argument("--random_seed", type=int, default=3407, help="Random seed for reproducibility.")
# other configs # other configs
parser.add_argument("--multi_turn", type=bool, default=False, help="Whether to use multi-turn convsersation training.") parser.add_argument("--window_size", type=int, default=None, help="the max length of the input sequence.")
parser.add_argument("--stride", type=int, default=None, help="the step size of the input sequence.")
parser.add_argument("--start_epoch", type=int, default=0, help="Start epoch for training.")
parser.add_argument("--start_batch", type=int, default=0, help="Start batch for training.")
parser.add_argument("--resume_from_checkpoint", type=bool, default=False, help="train from checkpoint or not.")
parser.add_argument("--dpo_beta", type=float, default=0.1, help="DPO beta value.") parser.add_argument("--dpo_beta", type=float, default=0.1, help="DPO beta value.")
args = parser.parse_args() args = parser.parse_args()
train( train(**vars(args))
param_path=args.param_path,
data_root_path=args.data_root_path,
n_epoch=args.n_epoch,
batch_size=args.batch_size,
accumulation_steps=args.accumulation_steps,
warmup_steps=args.warmup_steps,
max_lr=args.max_lr,
dpo_beta=args.dpo_beta,
adamw_betas=args.adamw_betas,
adamw_weight_decay=args.adamw_weight_decay,
max_grad_norm=args.max_grad_norm,
embdeding_lr_rate=args.embdeding_lr_rate,
checkpoint_interval=args.checkpoint_interval,
checkpoint_dir=args.checkpoint_dir,
train_type=args.train_type,
random_seed=args.random_seed,
multi_turn=args.multi_turn
)