Compare commits
26
Commits
82a3f2626f
...
v1.3.7
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a3275423a4 | ||
|
|
b37c3d000c | ||
|
|
6031020e37 | ||
|
|
c424dfc293 | ||
|
|
3a28e52e98 | ||
|
|
e371908b54 | ||
|
|
7c99da155c | ||
|
|
629e72385b | ||
|
|
0a708fff24 | ||
|
|
6e150ea6d0 | ||
|
|
cb8dcb97ea | ||
|
|
2d5dc93b3d | ||
|
|
4145d35e3c | ||
|
|
34c6c45bd6 | ||
|
|
e9def84ce7 | ||
|
|
836e02a166 | ||
|
|
b558e61f63 | ||
|
|
65ab69543b | ||
|
|
1d26aa2e93 | ||
|
|
a548d4553e | ||
|
|
dd1b39f435 | ||
|
|
94d6e713e9 | ||
|
|
47c37e4876 | ||
|
|
737585a32a | ||
|
|
a4688021bf | ||
|
|
7df6eb9211 |
+117
-106
@@ -22,7 +22,8 @@ classDiagram
|
|||||||
+int n_layers
|
+int n_layers
|
||||||
+float norm_eps
|
+float norm_eps
|
||||||
+int dim_ffn
|
+int dim_ffn
|
||||||
+bool tie_weight
|
+Optional[bool] tie_weight
|
||||||
|
+Optional[dict] rope_scaling
|
||||||
+int max_len
|
+int max_len
|
||||||
+float rope_theta
|
+float rope_theta
|
||||||
+str attn_type
|
+str attn_type
|
||||||
@@ -52,6 +53,7 @@ classDiagram
|
|||||||
+int n_kv_heads
|
+int n_kv_heads
|
||||||
+bool use_qk_norm
|
+bool use_qk_norm
|
||||||
+bool use_gated_attention
|
+bool use_gated_attention
|
||||||
|
+Optional[dict] rope_scaling
|
||||||
+Optional[str] pooling_type
|
+Optional[str] pooling_type
|
||||||
+Optional[bool] normalize_embeddings
|
+Optional[bool] normalize_embeddings
|
||||||
}
|
}
|
||||||
@@ -63,7 +65,7 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class TrainConfig {
|
class TrainConfig {
|
||||||
+nn.Module model
|
+Callable[[], nn.Module] model_fn
|
||||||
+str strategy
|
+str strategy
|
||||||
+Dataset dataset
|
+Dataset dataset
|
||||||
+Callable optimizer_fn
|
+Callable optimizer_fn
|
||||||
@@ -80,6 +82,7 @@ classDiagram
|
|||||||
+str log_dir
|
+str log_dir
|
||||||
+int log_interval
|
+int log_interval
|
||||||
+List[str] metrics
|
+List[str] metrics
|
||||||
|
+Optional[LoRAConfig] lora
|
||||||
+int random_seed
|
+int random_seed
|
||||||
+int num_workers
|
+int num_workers
|
||||||
+Optional[int] prefetch_factor
|
+Optional[int] prefetch_factor
|
||||||
@@ -104,8 +107,8 @@ classDiagram
|
|||||||
class BaseDataset {
|
class BaseDataset {
|
||||||
+int window_size
|
+int window_size
|
||||||
+int stride
|
+int stride
|
||||||
+Optional[BaseStorage] storage
|
+Optional[Store] storage
|
||||||
+load(load_path, storage_type, tokenizer)
|
+load(load_path, storage_type)
|
||||||
+__getitem__(index)
|
+__getitem__(index)
|
||||||
+__len__()
|
+__len__()
|
||||||
}
|
}
|
||||||
@@ -126,38 +129,25 @@ classDiagram
|
|||||||
+__getitem__(index) Dict
|
+__getitem__(index) Dict
|
||||||
}
|
}
|
||||||
|
|
||||||
class BaseSegmentFetcher {
|
class Store {
|
||||||
+List[Tensor] segments
|
+Dict[str, List[Tensor]] _data
|
||||||
+List[int] cum_lengths
|
+Dict[str, List[int]] _cum
|
||||||
+int total_length
|
+int _length
|
||||||
+fetch_data(begin_idx, end_idx) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseStorage {
|
|
||||||
+MultiSegmentFetcher _fetcher
|
|
||||||
+keys (property)
|
+keys (property)
|
||||||
+load(load_path, tokenizer)
|
+load(path)
|
||||||
+fetch(begin, end, keys)
|
+fetch(begin, end, keys)
|
||||||
+__len__()
|
+__len__()
|
||||||
|
-_fetch_key(key, begin, end) Tensor
|
||||||
|
-_normalize(raw)
|
||||||
}
|
}
|
||||||
|
|
||||||
class H5Storage {
|
class H5Store {
|
||||||
+load(load_path, tokenizer)
|
+load(path)
|
||||||
+fetch(begin, end, keys) Dict
|
|
||||||
+keys() List
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class JSONStorage {
|
class MmapStore {
|
||||||
+load(load_path, tokenizer)
|
+List _mmap_refs
|
||||||
+fetch(begin, end, keys) Dict
|
+load(path)
|
||||||
+keys() List
|
|
||||||
}
|
|
||||||
|
|
||||||
class MultiSegmentFetcher {
|
|
||||||
+Dict multi_fetchers
|
|
||||||
+List multi_keys
|
|
||||||
+key_fetch(begin_idx, end_idx, keys) Dict
|
|
||||||
+fetch_data(begin_idx, end_idx) Dict
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class ResumableDistributedSampler {
|
class ResumableDistributedSampler {
|
||||||
@@ -165,17 +155,17 @@ classDiagram
|
|||||||
+int iter
|
+int iter
|
||||||
}
|
}
|
||||||
|
|
||||||
class StorageFactory {
|
class StoreFactory {
|
||||||
+Registry _registry
|
+Registry _registry
|
||||||
+register(name) decorator
|
+register(name) decorator
|
||||||
+create(storage_type) BaseStorage
|
+create(storage_type) Store
|
||||||
}
|
}
|
||||||
|
|
||||||
class DatasetFactory {
|
class DatasetFactory {
|
||||||
+Registry _registry
|
+Registry _registry
|
||||||
+register(name) decorator
|
+register(name) decorator
|
||||||
+create(train_type, window_size, stride) BaseDataset
|
+create(train_type, window_size, stride) BaseDataset
|
||||||
+load(train_type, load_path, window_size, stride, storage_type, tokenizer) BaseDataset
|
+load(train_type, load_path, window_size, stride, storage_type) BaseDataset
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -186,8 +176,9 @@ classDiagram
|
|||||||
+int iteration
|
+int iteration
|
||||||
+dict extra
|
+dict extra
|
||||||
+dict meta
|
+dict meta
|
||||||
|
+dict config
|
||||||
+save(save_dir)
|
+save(save_dir)
|
||||||
+load(save_dir) Checkpoint
|
+load(save_dir, broadcast) Checkpoint
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -195,8 +186,8 @@ classDiagram
|
|||||||
class AutoModel {
|
class AutoModel {
|
||||||
+BaseModelConfig config
|
+BaseModelConfig config
|
||||||
+Registry _registry
|
+Registry _registry
|
||||||
+register(model_type) decorator
|
+register(name) decorator
|
||||||
+get_component_class(model_type) Type
|
+get_component_class(name) Type
|
||||||
+from_pretrained(path, disable_random_init, strict) nn.Module
|
+from_pretrained(path, disable_random_init, strict) nn.Module
|
||||||
+save_pretrained(save_directory)
|
+save_pretrained(save_directory)
|
||||||
+to(*args, **kwargs) Self
|
+to(*args, **kwargs) Self
|
||||||
@@ -210,7 +201,7 @@ classDiagram
|
|||||||
+RMSNorm norm
|
+RMSNorm norm
|
||||||
+Linear lm_head
|
+Linear lm_head
|
||||||
+forward(input_ids, input_mask, paged_cache, position_ids) Dict[str, Tensor]
|
+forward(input_ids, input_mask, paged_cache, position_ids) Dict[str, Tensor]
|
||||||
+load_state_dict(state_dict)
|
+load_state_dict(state_dict, strict, assign)
|
||||||
+state_dict()
|
+state_dict()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -235,6 +226,7 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class GQA {
|
class GQA {
|
||||||
|
+int dim
|
||||||
+int n_heads
|
+int n_heads
|
||||||
+int n_kv_heads
|
+int n_kv_heads
|
||||||
+int head_dim
|
+int head_dim
|
||||||
@@ -249,6 +241,7 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class MLA {
|
class MLA {
|
||||||
|
+int dim
|
||||||
+int n_heads
|
+int n_heads
|
||||||
+int n_kv_heads
|
+int n_kv_heads
|
||||||
+int head_dim
|
+int head_dim
|
||||||
@@ -309,6 +302,7 @@ classDiagram
|
|||||||
+int dim
|
+int dim
|
||||||
+int max_len
|
+int max_len
|
||||||
+float base
|
+float base
|
||||||
|
+Optional[Dict] rope_scaling
|
||||||
+forward(x, position_ids=None) Tensor
|
+forward(x, position_ids=None) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -321,10 +315,10 @@ classDiagram
|
|||||||
namespace tokenize {
|
namespace tokenize {
|
||||||
class AutoTokenizer {
|
class AutoTokenizer {
|
||||||
+vocab_size int
|
+vocab_size int
|
||||||
+encode(tokens, out_ids, add_special_tokens) List[int]
|
+encode(tokens, out_ids, is_pretokenized, add_special_tokens) List[int]
|
||||||
+decode(tokens, skip_special_tokens) str
|
+decode(tokens, skip_special_tokens) str
|
||||||
+__getattr__(name) Any (bos_id, eos_id, pad_id, stop_ids)
|
+__getattr__(name) Any (bos_id, eos_id, pad_id, stop_ids)
|
||||||
+apply_chat_template(messages, tokenize) Union[str, List[int]]
|
+apply_chat_template(messages, system_prompt, tokenize, add_generation_prompt) Union[str, List[int]]
|
||||||
+set_chat_template(template)
|
+set_chat_template(template)
|
||||||
+load(path)
|
+load(path)
|
||||||
+from_pretrained(path) AutoTokenizer
|
+from_pretrained(path) AutoTokenizer
|
||||||
@@ -332,7 +326,7 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class ChatTemplate {
|
class ChatTemplate {
|
||||||
+String template_str
|
+str template_str
|
||||||
+render(messages, system_prompt, **extra_variables) str
|
+render(messages, system_prompt, **extra_variables) str
|
||||||
+from_string(template) ChatTemplate
|
+from_string(template) ChatTemplate
|
||||||
}
|
}
|
||||||
@@ -370,6 +364,7 @@ classDiagram
|
|||||||
+SchedulerProtocol scheduler
|
+SchedulerProtocol scheduler
|
||||||
+Checkpoint checkpoint
|
+Checkpoint checkpoint
|
||||||
+TrainConfig config
|
+TrainConfig config
|
||||||
|
+dict model_config
|
||||||
+BaseExecutor executor
|
+BaseExecutor executor
|
||||||
+int epoch
|
+int epoch
|
||||||
+int iteration
|
+int iteration
|
||||||
@@ -383,7 +378,7 @@ classDiagram
|
|||||||
|
|
||||||
class TrainContextBuilder {
|
class TrainContextBuilder {
|
||||||
+TrainConfig config
|
+TrainConfig config
|
||||||
+with_checkpoint(checkpoint) TrainContextBuilder
|
+with_resume_dir(resume_dir) TrainContextBuilder
|
||||||
+build() TrainContext
|
+build() TrainContext
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -457,16 +452,15 @@ classDiagram
|
|||||||
+on_train_end(context)
|
+on_train_end(context)
|
||||||
+on_epoch_begin(context)
|
+on_epoch_begin(context)
|
||||||
+on_epoch_end(context)
|
+on_epoch_end(context)
|
||||||
+on_step_begin(context)
|
|
||||||
+on_step_end(context)
|
|
||||||
+on_batch_begin(context)
|
+on_batch_begin(context)
|
||||||
+on_batch_end(context)
|
+on_batch_end(context)
|
||||||
|
+on_optimizer_step(context)
|
||||||
+on_error(context)
|
+on_error(context)
|
||||||
}
|
}
|
||||||
|
|
||||||
class GradientClippingCallback {
|
class GradientClippingCallback {
|
||||||
+float max_grad_norm
|
+float max_grad_norm
|
||||||
+on_step_begin(context)
|
+on_optimizer_step(context)
|
||||||
}
|
}
|
||||||
|
|
||||||
class GradientCheckpointingCallback {
|
class GradientCheckpointingCallback {
|
||||||
@@ -479,16 +473,12 @@ classDiagram
|
|||||||
+str save_dir
|
+str save_dir
|
||||||
+int interval
|
+int interval
|
||||||
+bool weight_only
|
+bool weight_only
|
||||||
+Callable state_dict_fn
|
|
||||||
+Callable save_extra_fn
|
+Callable save_extra_fn
|
||||||
+Callable load_extra_fn
|
|
||||||
+_save_checkpoint(context)
|
+_save_checkpoint(context)
|
||||||
+on_train_begin(context)
|
|
||||||
+on_batch_end(context)
|
+on_batch_end(context)
|
||||||
+on_train_end(context)
|
+on_train_end(context)
|
||||||
+on_error(context)
|
+on_error(context)
|
||||||
+save_extra(context)$
|
+save_extra(context)$
|
||||||
+load_extra(extra, context)$
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class ProgressBarCallback {
|
class ProgressBarCallback {
|
||||||
@@ -512,7 +502,7 @@ classDiagram
|
|||||||
|
|
||||||
class ValidationCallback {
|
class ValidationCallback {
|
||||||
+_run_validation(context)
|
+_run_validation(context)
|
||||||
+on_step_end(context)
|
+on_optimizer_step(context)
|
||||||
}
|
}
|
||||||
|
|
||||||
class CallbackFactory {
|
class CallbackFactory {
|
||||||
@@ -525,7 +515,12 @@ classDiagram
|
|||||||
+float lr
|
+float lr
|
||||||
+float momentum
|
+float momentum
|
||||||
+float weight_decay
|
+float weight_decay
|
||||||
|
+bool nesterov
|
||||||
+int ns_steps
|
+int ns_steps
|
||||||
|
+float adamw_lr
|
||||||
|
+tuple adamw_betas
|
||||||
|
+float adamw_eps
|
||||||
|
+float adamw_wd
|
||||||
+step(closure) Optional[float]
|
+step(closure) Optional[float]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -546,6 +541,8 @@ classDiagram
|
|||||||
+AutoModel model
|
+AutoModel model
|
||||||
+AutoTokenizer tokenizer
|
+AutoTokenizer tokenizer
|
||||||
+KVCache page_cache
|
+KVCache page_cache
|
||||||
|
+Optional[str] device
|
||||||
|
+Optional[torch.dtype] dtype
|
||||||
+execute_prefill(tasks, prompt_len, start_pos)
|
+execute_prefill(tasks, prompt_len, start_pos)
|
||||||
+execute_decode(tasks) List[int]
|
+execute_decode(tasks) List[int]
|
||||||
}
|
}
|
||||||
@@ -557,7 +554,9 @@ classDiagram
|
|||||||
+bool _running
|
+bool _running
|
||||||
+Thread _loop_thread
|
+Thread _loop_thread
|
||||||
+int max_seq_len
|
+int max_seq_len
|
||||||
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
+str device
|
||||||
|
+torch.dtype dtype
|
||||||
|
+add_task(prompt, **kwargs) str
|
||||||
+remove_task(task_id)
|
+remove_task(task_id)
|
||||||
+start()
|
+start()
|
||||||
+stop()
|
+stop()
|
||||||
@@ -660,15 +659,19 @@ classDiagram
|
|||||||
|
|
||||||
class TaskManager {
|
class TaskManager {
|
||||||
+AutoTokenizer tokenizer
|
+AutoTokenizer tokenizer
|
||||||
|
+int max_batch_size
|
||||||
|
+int max_seq_len
|
||||||
|
+int max_prompt_len
|
||||||
+Deque waiting_queue
|
+Deque waiting_queue
|
||||||
+List active_tasks
|
+List active_tasks
|
||||||
+add_task(prompt, **kwargs) str
|
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
||||||
+remove_task(task_id) List[Task]
|
+remove_task(task_id) List[Task]
|
||||||
+remove_finished_tasks(stop_ids) List[Task]
|
+remove_finished_tasks(stop_ids) List[Task]
|
||||||
+pull_candidates(n) List[Task]
|
+pull_candidates(n) List[Task]
|
||||||
+activate(task)
|
+activate(task)
|
||||||
+return_to_waiting(tasks)
|
+return_to_waiting(tasks)
|
||||||
+get_active_tasks() List[Task]
|
+get_active_tasks() List[Task]
|
||||||
|
+get_stats() Dict
|
||||||
}
|
}
|
||||||
|
|
||||||
class GenerationRequest {
|
class GenerationRequest {
|
||||||
@@ -747,56 +750,58 @@ classDiagram
|
|||||||
+str model
|
+str model
|
||||||
+List[AnthropicMessage] messages
|
+List[AnthropicMessage] messages
|
||||||
+Optional[str] system
|
+Optional[str] system
|
||||||
+float temperature
|
+Optional[float] temperature
|
||||||
+float top_p
|
+Optional[float] top_p
|
||||||
+int top_k
|
+Optional[int] top_k
|
||||||
+int max_tokens
|
+int max_tokens
|
||||||
+bool stream
|
+Optional[bool] stream
|
||||||
+Optional[List[str]] stop_sequences
|
+Optional[List[str]] stop_sequences
|
||||||
}
|
}
|
||||||
|
|
||||||
class ProtocolHandler {
|
class ResponseBuilder {
|
||||||
<<abstract>>
|
<<abstract>>
|
||||||
|
+prepare(request, engine) Tuple[str, GenContext, List[str]]
|
||||||
|
+format_stream_start(ctx) List[str]
|
||||||
|
+format_chunk(token) str
|
||||||
|
+format_stream_end(ctx, stop) List[str]
|
||||||
|
+format_response(ctx, content, stop) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class OpenAIResponseBuilder {
|
||||||
|
+prepare(request, engine) Tuple
|
||||||
|
+format_stream_start(ctx) List[str]
|
||||||
|
+format_chunk(token) str
|
||||||
|
+format_stream_end(ctx, stop) List[str]
|
||||||
|
+format_response(ctx, content, stop) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class AnthropicResponseBuilder {
|
||||||
|
+prepare(request, engine) Tuple
|
||||||
|
+format_stream_start(ctx) List[str]
|
||||||
|
+format_chunk(token) str
|
||||||
|
+format_stream_end(ctx, stop) List[str]
|
||||||
|
+format_response(ctx, content, stop) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class ProtocolHandler {
|
||||||
+request
|
+request
|
||||||
+engine
|
+engine
|
||||||
+build_prompt() str
|
+builder: ResponseBuilder
|
||||||
+create_response_id() str
|
|
||||||
+get_stop_sequences() List[str]
|
|
||||||
+create_stop_checker() StopChecker
|
|
||||||
+on_token(ctx, token, stop_checker) Optional[str]
|
|
||||||
+format_stream_start(ctx) List[str]
|
|
||||||
+format_stream_token(ctx, token) str
|
|
||||||
+format_stream_end(ctx) List[str]
|
|
||||||
+format_non_stream_response(ctx, content) Dict
|
|
||||||
+handle() Union[StreamingResponse, Dict]
|
+handle() Union[StreamingResponse, Dict]
|
||||||
}
|
-_handle_stream(agen, ctx, stops) StreamingResponse
|
||||||
|
-_handle_non_stream(agen, ctx, stops) Dict
|
||||||
class OpenAIHandler {
|
|
||||||
+build_prompt() str
|
|
||||||
+create_response_id() str
|
|
||||||
}
|
|
||||||
|
|
||||||
class AnthropicHandler {
|
|
||||||
+build_prompt() str
|
|
||||||
+create_response_id() str
|
|
||||||
+on_token(ctx, token, stop_checker) Optional[str]
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class StopChecker {
|
class StopChecker {
|
||||||
+has_sequences (property) bool
|
|
||||||
+check(text) Optional[str]
|
+check(text) Optional[str]
|
||||||
+trim(text, matched) str
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class StreamContext {
|
class GenContext {
|
||||||
+str resp_id
|
+str resp_id
|
||||||
+int created
|
+int created
|
||||||
+str model
|
+str model
|
||||||
+int prompt_tokens
|
+int prompt_tokens
|
||||||
+int completion_tokens
|
+int completion_tokens
|
||||||
+str accumulated
|
|
||||||
+Optional[str] stop_matched
|
|
||||||
+str last_yield_trimmed
|
|
||||||
}
|
}
|
||||||
|
|
||||||
class app {
|
class app {
|
||||||
@@ -876,6 +881,11 @@ classDiagram
|
|||||||
+unwrap_model(model) nn.Module
|
+unwrap_model(model) nn.Module
|
||||||
}
|
}
|
||||||
|
|
||||||
|
class FSDPExecutor {
|
||||||
|
+_prepare_model(model) nn.Module
|
||||||
|
+unwrap_model(model) nn.Module
|
||||||
|
}
|
||||||
|
|
||||||
class ExecutorFactory {
|
class ExecutorFactory {
|
||||||
+Registry _registry
|
+Registry _registry
|
||||||
+register(name) decorator
|
+register(name) decorator
|
||||||
@@ -911,12 +921,13 @@ classDiagram
|
|||||||
TrainCallback <|-- CheckpointCallback
|
TrainCallback <|-- CheckpointCallback
|
||||||
TrainCallback <|-- ProgressBarCallback
|
TrainCallback <|-- ProgressBarCallback
|
||||||
TrainCallback <|-- MetricLoggerCallback
|
TrainCallback <|-- MetricLoggerCallback
|
||||||
|
TrainCallback <|-- ValidationCallback
|
||||||
BaseDataset <|-- SEQDataset
|
BaseDataset <|-- SEQDataset
|
||||||
BaseDataset <|-- SFTDataset
|
BaseDataset <|-- SFTDataset
|
||||||
BaseDataset <|-- DPODataset
|
BaseDataset <|-- DPODataset
|
||||||
BaseDataset <|-- GRPODataset
|
BaseDataset <|-- GRPODataset
|
||||||
BaseStorage <|-- H5Storage
|
Store <|-- H5Store
|
||||||
BaseStorage <|-- JSONStorage
|
Store <|-- MmapStore
|
||||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||||
BaseSamplingStrategy <|-- TopKStrategy
|
BaseSamplingStrategy <|-- TopKStrategy
|
||||||
BaseSamplingStrategy <|-- TopPStrategy
|
BaseSamplingStrategy <|-- TopPStrategy
|
||||||
@@ -936,20 +947,19 @@ classDiagram
|
|||||||
BaseFactory <|-- StrategyFactory
|
BaseFactory <|-- StrategyFactory
|
||||||
BaseFactory <|-- SchedulerFactory
|
BaseFactory <|-- SchedulerFactory
|
||||||
BaseFactory <|-- CallbackFactory
|
BaseFactory <|-- CallbackFactory
|
||||||
BaseFactory <|-- StorageFactory
|
BaseFactory <|-- StoreFactory
|
||||||
BaseFactory <|-- ExecutorFactory
|
BaseFactory <|-- ExecutorFactory
|
||||||
BaseFactory <|-- ConfigFactory
|
BaseFactory <|-- ConfigFactory
|
||||||
BaseExecutor <|-- NoneExecutor
|
BaseExecutor <|-- NoneExecutor
|
||||||
BaseExecutor <|-- DDPExecutor
|
BaseExecutor <|-- DDPExecutor
|
||||||
ProtocolHandler <|-- OpenAIHandler
|
BaseExecutor <|-- FSDPExecutor
|
||||||
ProtocolHandler <|-- AnthropicHandler
|
ResponseBuilder <|-- OpenAIResponseBuilder
|
||||||
|
ResponseBuilder <|-- AnthropicResponseBuilder
|
||||||
|
|
||||||
%% --- Composition (strong ownership, part destroyed with whole) ---
|
%% --- Composition (strong ownership, part destroyed with whole) ---
|
||||||
KVCache *-- PagePool
|
KVCache *-- PagePool
|
||||||
KVCache *-- Storage
|
KVCache *-- Storage
|
||||||
KVCache *-- TaskTable
|
KVCache *-- TaskTable
|
||||||
PagePool *-- Allocator
|
|
||||||
PagePool *-- PrefixCache
|
|
||||||
InferenceEngine *-- InferenceScheduler
|
InferenceEngine *-- InferenceScheduler
|
||||||
InferenceScheduler *-- KVCache
|
InferenceScheduler *-- KVCache
|
||||||
InferenceScheduler *-- Executor
|
InferenceScheduler *-- Executor
|
||||||
@@ -963,7 +973,6 @@ classDiagram
|
|||||||
DecoderBlock *-- RMSNorm
|
DecoderBlock *-- RMSNorm
|
||||||
ChatCompletionRequest *-- ChatMessage
|
ChatCompletionRequest *-- ChatMessage
|
||||||
MessagesRequest *-- AnthropicMessage
|
MessagesRequest *-- AnthropicMessage
|
||||||
AutoTokenizer *-- ChatTemplate
|
|
||||||
BaseFactory *-- Registry
|
BaseFactory *-- Registry
|
||||||
BaseExecutor *-- GradientState
|
BaseExecutor *-- GradientState
|
||||||
AccumOptimizer o-- GradientState
|
AccumOptimizer o-- GradientState
|
||||||
@@ -971,6 +980,9 @@ classDiagram
|
|||||||
|
|
||||||
%% --- Aggregation (weak ownership) ---
|
%% --- Aggregation (weak ownership) ---
|
||||||
AutoModel o-- BaseModelConfig
|
AutoModel o-- BaseModelConfig
|
||||||
|
AutoTokenizer o-- ChatTemplate
|
||||||
|
PagePool o-- Allocator
|
||||||
|
PagePool o-- PrefixCache
|
||||||
Trainer o-- TrainCallback
|
Trainer o-- TrainCallback
|
||||||
TrainContext o-- BaseStrategy
|
TrainContext o-- BaseStrategy
|
||||||
TrainContext o-- BaseScheduler
|
TrainContext o-- BaseScheduler
|
||||||
@@ -978,7 +990,7 @@ classDiagram
|
|||||||
TrainContext o-- BaseExecutor
|
TrainContext o-- BaseExecutor
|
||||||
KvcacheView o-- Storage
|
KvcacheView o-- Storage
|
||||||
SamplingPipeline o-- BaseSamplingStrategy
|
SamplingPipeline o-- BaseSamplingStrategy
|
||||||
BaseDataset o-- BaseStorage
|
BaseDataset o-- Store
|
||||||
|
|
||||||
%% --- Dependency (uses temporarily) ---
|
%% --- Dependency (uses temporarily) ---
|
||||||
TrainConfig ..> BaseStrategy : selects
|
TrainConfig ..> BaseStrategy : selects
|
||||||
@@ -992,12 +1004,13 @@ classDiagram
|
|||||||
FFNFactory ..> DeepSeekMoE : creates
|
FFNFactory ..> DeepSeekMoE : creates
|
||||||
DecoderBlock ..> AttnFactory : uses
|
DecoderBlock ..> AttnFactory : uses
|
||||||
DecoderBlock ..> FFNFactory : uses
|
DecoderBlock ..> FFNFactory : uses
|
||||||
StorageFactory ..> H5Storage : creates
|
StoreFactory ..> H5Store : creates
|
||||||
StorageFactory ..> JSONStorage : creates
|
StoreFactory ..> MmapStore : creates
|
||||||
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
||||||
ConfigFactory ..> EncoderConfig : creates
|
ConfigFactory ..> EncoderConfig : creates
|
||||||
ExecutorFactory ..> NoneExecutor : creates
|
ExecutorFactory ..> NoneExecutor : creates
|
||||||
ExecutorFactory ..> DDPExecutor : creates
|
ExecutorFactory ..> DDPExecutor : creates
|
||||||
|
ExecutorFactory ..> FSDPExecutor : creates
|
||||||
TrainContextBuilder ..> ExecutorFactory : creates
|
TrainContextBuilder ..> ExecutorFactory : creates
|
||||||
Trainer ..> TrainContextBuilder : uses
|
Trainer ..> TrainContextBuilder : uses
|
||||||
TrainContextBuilder ..> TrainContext : creates
|
TrainContextBuilder ..> TrainContext : creates
|
||||||
@@ -1009,10 +1022,10 @@ classDiagram
|
|||||||
KVCache ..> KvcacheView : binds
|
KVCache ..> KvcacheView : binds
|
||||||
InferenceEngine ..> GenerationRequest : uses
|
InferenceEngine ..> GenerationRequest : uses
|
||||||
InferenceEngine ..> GenerateResult : creates
|
InferenceEngine ..> GenerateResult : creates
|
||||||
OpenAIHandler ..> ChatCompletionRequest : receives
|
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
||||||
AnthropicHandler ..> MessagesRequest : receives
|
AnthropicResponseBuilder ..> MessagesRequest : receives
|
||||||
ProtocolHandler ..> StopChecker : creates
|
ProtocolHandler ..> StopChecker : creates
|
||||||
ProtocolHandler ..> StreamContext : creates
|
ProtocolHandler ..> GenContext : creates
|
||||||
|
|
||||||
%% --- Association (general usage) ---
|
%% --- Association (general usage) ---
|
||||||
Trainer --> TrainConfig
|
Trainer --> TrainConfig
|
||||||
@@ -1025,8 +1038,6 @@ classDiagram
|
|||||||
Executor --> AutoModel
|
Executor --> AutoModel
|
||||||
Executor --> AutoTokenizer
|
Executor --> AutoTokenizer
|
||||||
TaskManager --> AutoTokenizer
|
TaskManager --> AutoTokenizer
|
||||||
MultiSegmentFetcher --> BaseSegmentFetcher
|
|
||||||
ResumableDistributedSampler --> BaseDataset
|
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -1036,13 +1047,13 @@ classDiagram
|
|||||||
| Module | Components | Description |
|
| Module | Components | Description |
|
||||||
|--------|------------|-------------|
|
|--------|------------|-------------|
|
||||||
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
||||||
| **astrai.dataset** | BaseDataset–GRPODataset, BaseStorage–JSONStorage, StorageFactory, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
| **astrai.dataset** | BaseDataset–GRPODataset, Store–MmapStore, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
||||||
| **astrai.serialization** | Checkpoint | Model serialization |
|
| **astrai.serialization** | Checkpoint | Model serialization |
|
||||||
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
||||||
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–SGDRScheduler, SchedulerFactory, TrainCallback(Protocol)–ValidationCallback, CallbackFactory, Muon | Training workflow |
|
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–SGDRScheduler, SchedulerFactory, TrainCallback(Protocol)–ValidationCallback, CallbackFactory, Muon | Training workflow |
|
||||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–KvcacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler–AnthropicHandler, StopChecker, StreamContext, ChatMessage–MessagesRequest, app | Inference service |
|
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–KvcacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessage–MessagesRequest, app | Inference service |
|
||||||
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation |
|
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation |
|
||||||
| **astrai.factory** | Registry, BaseFactory[T] | Component registration |
|
| **astrai.factory** | Registry, BaseFactory[T] | Component registration |
|
||||||
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
||||||
|
|
||||||
@@ -1050,17 +1061,17 @@ classDiagram
|
|||||||
|
|
||||||
| Pattern | Classes | Purpose |
|
| Pattern | Classes | Purpose |
|
||||||
|---------|---------|---------|
|
|---------|---------|---------|
|
||||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StorageFactory`, `ConfigFactory`, `ExecutorFactory` | Decorator-based component creation |
|
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory` | Decorator-based component creation |
|
||||||
| **Registry** | `BaseFactory`, `Registry` | Component registration with category/priority |
|
| **Registry** | `BaseFactory`, `Registry` | Component registration with category/priority |
|
||||||
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
||||||
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
||||||
| **Template Method** | `ProtocolHandler`, `OpenAIHandler`, `AnthropicHandler` | HTTP API handler with format hooks |
|
| **Strategy (API)** | `ResponseBuilder`, `OpenAIResponseBuilder`, `AnthropicResponseBuilder` | HTTP API handler with format hooks |
|
||||||
| **Builder** | `TrainContextBuilder` | Chain-building training context |
|
| **Builder** | `TrainContextBuilder` | Chain-building training context |
|
||||||
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
||||||
| **Context** | `TrainContext` | Unified training state bag |
|
| **Context** | `TrainContext` | Unified training state bag |
|
||||||
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
||||||
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor` | Gradient accumulation & model distribution |
|
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor` | Gradient accumulation & model distribution |
|
||||||
| **Storage** | `BaseStorage`, `H5Storage`, `JSONStorage` | Format-agnostic data access |
|
| **Storage** | `Store`, `H5Store`, `MmapStore` | Format-agnostic data access with multi-segment support |
|
||||||
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
||||||
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
||||||
|
|
||||||
@@ -1069,13 +1080,13 @@ classDiagram
|
|||||||
1. **Config → Training**: `TrainConfig` holds model, dataset, optimizer_fn, scheduler_fn, `parallel_mode`, `executor_kwargs`
|
1. **Config → Training**: `TrainConfig` holds model, dataset, optimizer_fn, scheduler_fn, `parallel_mode`, `executor_kwargs`
|
||||||
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution
|
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution
|
||||||
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
|
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
|
||||||
4. **Executor Selection**: `ExecutorFactory.create(parallel_mode, **executor_kwargs)` → `NoneExecutor` (single) / `DDPExecutor` (distributed)
|
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
|
||||||
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
|
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
|
||||||
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
||||||
7. **Dataset Loading**: `DatasetFactory` creates datasets, `BaseStorage` (H5Storage/JSONStorage) loads via `BaseSegmentFetcher` + `MultiSegmentFetcher`
|
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore) loads data with explicit `_length` and multi-segment `_data`
|
||||||
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
||||||
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`
|
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`
|
||||||
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
||||||
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
|
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
|
||||||
|
|
||||||
> Document Update Time: 2026-05-24
|
> Document Update Time: 2026-05-28
|
||||||
|
|||||||
+10
-10
@@ -5,21 +5,21 @@ This document describes the data pipeline: from raw text to model input tensors.
|
|||||||
## Overview
|
## Overview
|
||||||
|
|
||||||
```
|
```
|
||||||
Raw Text → AutoTokenizer → Token IDs → .h5/.json → Dataset → Sampler → DataLoader → Training/Inference
|
Raw Text → AutoTokenizer → Token IDs → .h5/.bin → Dataset → Sampler → DataLoader → Training/Inference
|
||||||
```
|
```
|
||||||
|
|
||||||
## Data Preparation
|
## Data Preparation
|
||||||
|
|
||||||
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or JSON (`.json`/`.jsonl`) files with keyed tensor groups.
|
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or binary (`.bin` + `meta.json`) files with keyed tensor groups.
|
||||||
|
|
||||||
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
||||||
|
|
||||||
```
|
```
|
||||||
StorageFactory.create("h5") → H5Storage
|
StoreFactory.create("h5") → H5Store
|
||||||
StorageFactory.create("json") → JSONStorage
|
StoreFactory.create("bin") → MmapStore
|
||||||
```
|
```
|
||||||
|
|
||||||
Both support shared memory via `.share_memory_()`.
|
H5 backend supports shared memory via `.share_memory_()`. Bin (mmap) uses OS page-cache sharing natively.
|
||||||
|
|
||||||
## Data Keys by Training Type
|
## Data Keys by Training Type
|
||||||
|
|
||||||
@@ -33,14 +33,14 @@ Both support shared memory via `.share_memory_()`.
|
|||||||
## Dataset Architecture
|
## Dataset Architecture
|
||||||
|
|
||||||
```
|
```
|
||||||
DatasetFactory.load(train_type, path, window_size, stride)
|
DatasetFactory.load(train_type, load_path, window_size, stride, storage_type)
|
||||||
→ StorageFactory.create(detect_format(path))
|
→ StoreFactory.create(detect_format(path))
|
||||||
→ MultiSegmentFetcher(BaseSegmentFetcher per key)
|
→ Store._data[Dict[str, List[Tensor]]] + _cum[Dict[str, List[int]]]
|
||||||
→ BaseDataset.__getitem__(idx)
|
→ BaseDataset.__getitem__(idx)
|
||||||
→ sliding window [begin, end) via get_index(idx)
|
→ sliding window [begin, end) via get_index(idx)
|
||||||
```
|
```
|
||||||
|
|
||||||
`window_size` = max input length, `stride` = step between consecutive samples.
|
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`).
|
||||||
|
|
||||||
## Sampler
|
## Sampler
|
||||||
|
|
||||||
@@ -54,4 +54,4 @@ DatasetFactory.load(train_type, path, window_size, stride)
|
|||||||
|
|
||||||
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
||||||
|
|
||||||
> Document Update Time: 2026-05-17
|
> Document Update Time: 2026-05-28
|
||||||
|
|||||||
+26
-18
@@ -16,12 +16,12 @@ Six classes working together:
|
|||||||
|
|
||||||
```
|
```
|
||||||
KVCache (facade)
|
KVCache (facade)
|
||||||
├── Allocator bitmask-based page allocator + ref-count + LRU eviction
|
├── PagePool orchestrates page allocation + prefix matching
|
||||||
├── PrefixCache hash-based prefix matching (page_hash via rolling hash)
|
│ ├── Allocator bitmask-based page allocator + ref-count + LRU eviction (inside PagePool)
|
||||||
├── PagePool orchestrates Allocator + PrefixCache
|
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash) (inside PagePool)
|
||||||
├── TaskTable maps task_id → page_table + cached token count
|
├── TaskTable maps task_id → page_table + cached token count
|
||||||
├── Storage k_cache / v_cache tensors (n_layers × n_pages × page_size × n_kv_heads × head_dim)
|
├── Storage k_cache / v_cache tensors (n_layers × n_pages × page_size × n_kv_heads × head_dim)
|
||||||
└── KvcacheView bundles Storage + page_table + total_len for attention layers
|
└── KvcacheView bundles Storage + page_table + total_len for attention layers (returned by bind())
|
||||||
```
|
```
|
||||||
|
|
||||||
`KVCache.bind(page_table, total_len)` returns a `KvcacheView` used by attention layers via `write()` / `gather()`.
|
`KVCache.bind(page_table, total_len)` returns a `KvcacheView` used by attention layers via `write()` / `gather()`.
|
||||||
@@ -40,26 +40,32 @@ KVCache (facade)
|
|||||||
## Sampling (Strategy Pattern)
|
## Sampling (Strategy Pattern)
|
||||||
|
|
||||||
```
|
```
|
||||||
BaseSamplingStrategy → TemperatureStrategy → TopKStrategy → TopPStrategy
|
BaseSamplingStrategy (ABC)
|
||||||
|
├── TemperatureStrategy
|
||||||
|
├── TopKStrategy
|
||||||
|
└── TopPStrategy
|
||||||
```
|
```
|
||||||
|
|
||||||
`SamplingPipeline` composes them: Temperature → Top-K → Top-P → softmax → multinomial.
|
`SamplingPipeline` composes them: Temperature → Top-K → Top-P → softmax → multinomial.
|
||||||
`sample()` is a convenience shortcut for one-shot usage.
|
`sample()` is a convenience shortcut for one-shot usage.
|
||||||
|
|
||||||
## Protocol Handlers (Template Method)
|
## Protocol Handlers (Strategy Pattern)
|
||||||
|
|
||||||
```python
|
```python
|
||||||
class ProtocolHandler(ABC):
|
class ProtocolHandler: # concrete orchestrator
|
||||||
def handle(self):
|
def __init__(self, request, engine, builder): ...
|
||||||
ctx = StreamContext(...)
|
async def handle(self):
|
||||||
|
prompt, ctx, stops = builder.prepare(request, engine)
|
||||||
agen = engine.generate_async(prompt, ...)
|
agen = engine.generate_async(prompt, ...)
|
||||||
if stream: self._handle_stream(agen, ctx)
|
if stream: self._handle_stream(agen, ctx, stops)
|
||||||
else: self._handle_non_stream(agen, ctx)
|
else: return await self._handle_non_stream(agen, ctx, stops)
|
||||||
```
|
```
|
||||||
|
|
||||||
Subclass hooks: `build_prompt()`, `create_response_id()`, `format_stream_start/token/end()`, `format_non_stream_response()`.
|
`ResponseBuilder` (ABC): `prepare()`, `format_stream_start()`, `format_chunk()`, `format_stream_end()`, `format_response()`.
|
||||||
|
|
||||||
`OpenAIHandler` → `/v1/chat/completions`, `AnthropicHandler` → `/v1/messages`.
|
`OpenAIResponseBuilder` → `/v1/chat/completions`, `AnthropicResponseBuilder` → `/v1/messages`.
|
||||||
|
|
||||||
|
Adding a protocol = one builder file, no handler subclassing needed.
|
||||||
|
|
||||||
## Engine & GenerateResult
|
## Engine & GenerateResult
|
||||||
|
|
||||||
@@ -94,12 +100,14 @@ Response:
|
|||||||
{
|
{
|
||||||
"id": "chatcmpl-abc123",
|
"id": "chatcmpl-abc123",
|
||||||
"object": "chat.completion",
|
"object": "chat.completion",
|
||||||
"choices": [{"message": {"role": "assistant", "content": "Hello!"}, "finish_reason": "stop"}],
|
"created": 1717000000,
|
||||||
|
"model": "astrai",
|
||||||
|
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Hello!"}, "finish_reason": "stop"}],
|
||||||
"usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}
|
"usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
Streaming SSE: `data: {"choices":[{"delta":{"role":"assistant"}}]}` → token chunks → `data: [DONE]`
|
Streaming SSE: `object: "chat.completion.chunk"` — starts with role delta, then token chunks, ends with finish chunk + usage stats, then `data: [DONE]`.
|
||||||
|
|
||||||
### Anthropic
|
### Anthropic
|
||||||
|
|
||||||
@@ -116,10 +124,10 @@ Supports `stop_sequences` and streaming via `event: content_block_delta`.
|
|||||||
| Param | Type | Default | Description |
|
| Param | Type | Default | Description |
|
||||||
|-------|------|---------|-------------|
|
|-------|------|---------|-------------|
|
||||||
| `messages` | List[dict] | required | Chat messages (role, content) |
|
| `messages` | List[dict] | required | Chat messages (role, content) |
|
||||||
| `temperature` | float | 1.0 | Sampling temperature (0.0–2.0) |
|
| `temperature` | float | 1.0 | Sampling temperature (>= 0.0) |
|
||||||
| `top_p` | float | 1.0 | Nucleus threshold |
|
| `top_p` | float | 1.0 | Nucleus threshold |
|
||||||
| `top_k` | int | 50 | Top-k count |
|
| `top_k` | int | 50 | Top-k count |
|
||||||
| `max_tokens` | int | None | Max generation length |
|
| `max_tokens` | Optional[int] | None | Max generation length |
|
||||||
| `stream` | bool | False | Stream output |
|
| `stream` | bool | False | Stream output |
|
||||||
|
|
||||||
## Engine API
|
## Engine API
|
||||||
@@ -137,4 +145,4 @@ engine.generate(["A", "B"], stream=True) # -> Generator[Tuple[int, str]]
|
|||||||
await engine.generate_async("Hello", ...) # -> AsyncGenerator[str]
|
await engine.generate_async("Hello", ...) # -> AsyncGenerator[str]
|
||||||
```
|
```
|
||||||
|
|
||||||
> Document Update Time: 2026-05-17
|
> Document Update Time: 2026-05-28
|
||||||
|
|||||||
@@ -53,7 +53,7 @@
|
|||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--nprocs` | Number of GPUs / processes | 1 |
|
| `--nprocs` | Number of GPUs / processes | 1 |
|
||||||
| `--parallel_mode` | Parallel strategy (`none` or `ddp`) | none |
|
| `--parallel_mode` | Parallel strategy (`none`, `ddp`, or `fsdp`) | none |
|
||||||
| `--device_type` | Device type | cuda |
|
| `--device_type` | Device type | cuda |
|
||||||
| `--start_method` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | spawn |
|
| `--start_method` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | spawn |
|
||||||
|
|
||||||
|
|||||||
+10
-9
@@ -74,7 +74,8 @@ on_train_begin
|
|||||||
on_batch_begin
|
on_batch_begin
|
||||||
with executor.accumulate(model):
|
with executor.accumulate(model):
|
||||||
loss = strategy(batch)
|
loss = strategy(batch)
|
||||||
(loss / grad_accum_steps).backward()
|
stand_loss = loss / executor.grad_accum_steps
|
||||||
|
executor.backward(stand_loss)
|
||||||
iteration += 1
|
iteration += 1
|
||||||
on_batch_end
|
on_batch_end
|
||||||
|
|
||||||
@@ -82,8 +83,8 @@ on_train_begin
|
|||||||
on_optimizer_step
|
on_optimizer_step
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
|
if scheduler:
|
||||||
scheduler.step() # called every iteration
|
scheduler.step()
|
||||||
on_epoch_end
|
on_epoch_end
|
||||||
on_train_end
|
on_train_end
|
||||||
```
|
```
|
||||||
@@ -170,27 +171,27 @@ Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoi
|
|||||||
## Checkpoint
|
## Checkpoint
|
||||||
|
|
||||||
```
|
```
|
||||||
Checkpoint(state_dict, epoch, iteration, extra, meta)
|
Checkpoint(state_dict, epoch, iteration, extra, meta, config)
|
||||||
├── save(save_dir) rank-0 only: meta.json (includes training config) + state_dict.safetensors + optional optimizer.pt / scheduler.pt
|
├── save(save_dir) rank-0 only: meta.json (epoch/iteration/timestamp) + config.json (model config) + state_dict.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt)
|
||||||
└── load(save_dir) broadcasts metadata from rank-0
|
└── load(save_dir) broadcasts metadata from rank-0
|
||||||
```
|
```
|
||||||
|
|
||||||
Optimizer/scheduler state persisted by default via `Checkpoint.extra`.
|
Optimizer/scheduler state persisted by default via `Checkpoint.extra`.
|
||||||
Training config (`TrainConfig.to_dict()`) saved into `meta.json` during training via `CheckpointCallback`.
|
Model config (`context.model_config`) saved into `config.json` during training via `CheckpointCallback`.
|
||||||
|
|
||||||
## TrainContextBuilder (Builder Pattern)
|
## TrainContextBuilder (Builder Pattern)
|
||||||
|
|
||||||
```python
|
```python
|
||||||
context = (
|
context = (
|
||||||
TrainContextBuilder(config)
|
TrainContextBuilder(config)
|
||||||
.with_checkpoint(checkpoint)
|
.with_resume_dir(resume_dir)
|
||||||
.build()
|
.build()
|
||||||
)
|
)
|
||||||
# Returns TrainContext with model, strategy, optimizer, scheduler, dataloader, checkpoint
|
# Returns TrainContext with model, strategy, optimizer, scheduler, dataloader, checkpoint
|
||||||
```
|
```
|
||||||
|
|
||||||
- Loads checkpoint weights if provided
|
- Loads checkpoint weights if provided
|
||||||
- Creates executor via `ExecutorFactory.create(parallel_mode, **executor_kwargs)`
|
- Creates executor via `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)`
|
||||||
- Calls `executor.prepare(model, optimizer, dataloader, scheduler)` for model distribution (e.g. DDP) + gradient accumulation wrappers
|
- Calls `executor.prepare(model, optimizer, dataloader, scheduler)` for model distribution (e.g. DDP) + gradient accumulation wrappers
|
||||||
- Creates `ResumableDistributedSampler` for shuffle+resume
|
- Creates `ResumableDistributedSampler` for shuffle+resume
|
||||||
- Builds strategy via `StrategyFactory.create(train_type, ...)`
|
- Builds strategy via `StrategyFactory.create(train_type, ...)`
|
||||||
@@ -223,4 +224,4 @@ nohup python scripts/tools/train.py \
|
|||||||
|
|
||||||
Full parameter reference at [params.md](params.md).
|
Full parameter reference at [params.md](params.md).
|
||||||
|
|
||||||
> Document Update Time: 2026-05-24
|
> Document Update Time: 2026-05-28
|
||||||
|
|||||||
+1
-1
@@ -1,4 +1,4 @@
|
|||||||
__version__ = "1.3.6"
|
__version__ = "1.3.7"
|
||||||
__author__ = "ViperEkura"
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
from astrai.config import (
|
from astrai.config import (
|
||||||
|
|||||||
+12
-3
@@ -13,12 +13,21 @@ class BaseConfig:
|
|||||||
d[fld.name] = v
|
d[fld.name] = v
|
||||||
elif v is None:
|
elif v is None:
|
||||||
d[fld.name] = None
|
d[fld.name] = None
|
||||||
elif isinstance(v, (dict, list)):
|
elif isinstance(v, (dict, list, tuple)):
|
||||||
try:
|
try:
|
||||||
json.dumps(v)
|
val = list(v) if isinstance(v, tuple) else v
|
||||||
d[fld.name] = v
|
json.dumps(val)
|
||||||
|
d[fld.name] = val
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
pass
|
pass
|
||||||
|
elif isinstance(v, BaseConfig):
|
||||||
|
d[fld.name] = v.to_dict()
|
||||||
|
elif hasattr(v, "__dataclass_fields__"):
|
||||||
|
sub = {}
|
||||||
|
for f in fields(v):
|
||||||
|
a = getattr(v, f.name)
|
||||||
|
sub[f.name] = list(a) if isinstance(a, tuple) else a
|
||||||
|
d[fld.name] = sub
|
||||||
return d
|
return d
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
@@ -49,6 +49,7 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
|||||||
|
|
||||||
max_len: Optional[int] = None
|
max_len: Optional[int] = None
|
||||||
rope_theta: Optional[float] = None
|
rope_theta: Optional[float] = None
|
||||||
|
rope_scaling: Optional[dict] = None
|
||||||
|
|
||||||
attn_type: str = "gqa"
|
attn_type: str = "gqa"
|
||||||
n_heads: Optional[int] = None
|
n_heads: Optional[int] = None
|
||||||
@@ -80,6 +81,7 @@ class EncoderConfig(BaseModelConfig):
|
|||||||
|
|
||||||
max_len: Optional[int] = None
|
max_len: Optional[int] = None
|
||||||
rope_theta: Optional[float] = None
|
rope_theta: Optional[float] = None
|
||||||
|
rope_scaling: Optional[dict] = None
|
||||||
|
|
||||||
n_heads: Optional[int] = None
|
n_heads: Optional[int] = None
|
||||||
n_kv_heads: Optional[int] = None
|
n_kv_heads: Optional[int] = None
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from torch.optim.lr_scheduler import LRScheduler
|
|||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
from astrai.config.base import BaseConfig
|
from astrai.config.base import BaseConfig
|
||||||
|
from astrai.model.components.lora import LoRAConfig
|
||||||
|
|
||||||
|
|
||||||
def required(**kw):
|
def required(**kw):
|
||||||
@@ -16,8 +17,8 @@ def required(**kw):
|
|||||||
@dataclass
|
@dataclass
|
||||||
class TrainConfig(BaseConfig):
|
class TrainConfig(BaseConfig):
|
||||||
# basic setting
|
# basic setting
|
||||||
model: nn.Module = field(
|
model_fn: Callable[[], nn.Module] = field(
|
||||||
default=None, metadata=required(help="Model for training.")
|
default=None, metadata=required(help="Model factory for training.")
|
||||||
)
|
)
|
||||||
strategy: str = field(default=None, metadata=required(help="Training strategy."))
|
strategy: str = field(default=None, metadata=required(help="Training strategy."))
|
||||||
dataset: Dataset = field(
|
dataset: Dataset = field(
|
||||||
@@ -56,6 +57,12 @@ class TrainConfig(BaseConfig):
|
|||||||
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# lora setting
|
||||||
|
lora: Optional[LoRAConfig] = field(
|
||||||
|
default=None,
|
||||||
|
metadata={"help": "LoRA config. None means full fine-tuning."},
|
||||||
|
)
|
||||||
|
|
||||||
# metric setting
|
# metric setting
|
||||||
log_dir: str = field(
|
log_dir: str = field(
|
||||||
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
|
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
|
||||||
@@ -97,7 +104,7 @@ class TrainConfig(BaseConfig):
|
|||||||
)
|
)
|
||||||
parallel_mode: str = field(
|
parallel_mode: str = field(
|
||||||
default="none",
|
default="none",
|
||||||
metadata={"help": "Parallel strategy: none, ddp."},
|
metadata={"help": "Parallel strategy: none, ddp, fsdp."},
|
||||||
)
|
)
|
||||||
start_method: str = field(
|
start_method: str = field(
|
||||||
default="spawn",
|
default="spawn",
|
||||||
|
|||||||
+12
-16
@@ -4,32 +4,28 @@ from astrai.dataset.dataset import (
|
|||||||
)
|
)
|
||||||
from astrai.dataset.sampler import ResumableDistributedSampler
|
from astrai.dataset.sampler import ResumableDistributedSampler
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
BaseSegmentFetcher,
|
H5Store,
|
||||||
BaseStorage,
|
MmapStore,
|
||||||
H5Storage,
|
Store,
|
||||||
JSONStorage,
|
StoreFactory,
|
||||||
MultiSegmentFetcher,
|
|
||||||
StorageFactory,
|
|
||||||
detect_format,
|
detect_format,
|
||||||
|
load_bin,
|
||||||
load_h5,
|
load_h5,
|
||||||
load_json,
|
save_bin,
|
||||||
save_h5,
|
save_h5,
|
||||||
save_json,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"BaseDataset",
|
"BaseDataset",
|
||||||
"DatasetFactory",
|
"DatasetFactory",
|
||||||
"BaseSegmentFetcher",
|
"Store",
|
||||||
"MultiSegmentFetcher",
|
"StoreFactory",
|
||||||
"BaseStorage",
|
"H5Store",
|
||||||
"H5Storage",
|
"MmapStore",
|
||||||
"JSONStorage",
|
|
||||||
"StorageFactory",
|
|
||||||
"detect_format",
|
"detect_format",
|
||||||
"save_h5",
|
"save_h5",
|
||||||
"load_h5",
|
"load_h5",
|
||||||
"save_json",
|
"save_bin",
|
||||||
"load_json",
|
"load_bin",
|
||||||
"ResumableDistributedSampler",
|
"ResumableDistributedSampler",
|
||||||
]
|
]
|
||||||
|
|||||||
+15
-26
@@ -8,8 +8,8 @@ from torch import Tensor
|
|||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
BaseStorage,
|
Store,
|
||||||
StorageFactory,
|
StoreFactory,
|
||||||
detect_format,
|
detect_format,
|
||||||
)
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
@@ -26,7 +26,7 @@ class BaseDataset(Dataset, ABC):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.window_size = window_size
|
self.window_size = window_size
|
||||||
self.stride = stride
|
self.stride = stride
|
||||||
self.storage: Optional[BaseStorage] = None
|
self.storage: Optional[Store] = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def required_keys(self) -> List[str]:
|
def required_keys(self) -> List[str]:
|
||||||
@@ -48,37 +48,26 @@ class BaseDataset(Dataset, ABC):
|
|||||||
f"Missing: {missing}"
|
f"Missing: {missing}"
|
||||||
)
|
)
|
||||||
|
|
||||||
def load(self, load_path: str, storage_type: Optional[str] = None, tokenizer=None):
|
def load(self, load_path: str, storage_type: Optional[str] = None):
|
||||||
"""Load dataset from the given path.
|
"""Load dataset from the given path.
|
||||||
|
|
||||||
Auto-detects the storage format if not specified.
|
Auto-detects the storage format if not specified.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
load_path: Path to the data directory or file
|
load_path: Path to the data directory or file
|
||||||
storage_type: Force a specific storage type ("h5", "json"),
|
storage_type: Force a specific storage type ("h5", "bin"),
|
||||||
or None for auto-detection
|
or None for auto-detection
|
||||||
tokenizer: Callable str -> List[int], used to tokenize raw text
|
|
||||||
in JSON files. Ignored for HDF5.
|
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
KeyError: If the loaded storage is missing required keys.
|
KeyError: If the loaded storage is missing required keys.
|
||||||
"""
|
"""
|
||||||
if storage_type is None:
|
if storage_type is None:
|
||||||
storage_type = detect_format(load_path)
|
storage_type = detect_format(load_path)
|
||||||
self.storage = StorageFactory.create(storage_type)
|
self.storage = StoreFactory.create(storage_type)
|
||||||
self._load_path = load_path
|
self._load_path = load_path
|
||||||
self.storage.load(load_path, tokenizer=tokenizer)
|
self.storage.load(load_path)
|
||||||
self._validate_keys()
|
self._validate_keys()
|
||||||
|
|
||||||
def load_json(self, load_path: str, tokenizer=None):
|
|
||||||
"""Load dataset from JSON files explicitly.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
load_path: Path to the JSON data file or directory
|
|
||||||
tokenizer: Optional tokenizer callable for raw text JSON.
|
|
||||||
"""
|
|
||||||
self.load(load_path, storage_type="json", tokenizer=tokenizer)
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def count(self) -> int:
|
def count(self) -> int:
|
||||||
"""Return the total number of raw elements (tokens) in the dataset."""
|
"""Return the total number of raw elements (tokens) in the dataset."""
|
||||||
@@ -148,7 +137,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _validate_component(cls, dataset_cls: type) -> None:
|
def _validate_component(cls, dataset_cls: type):
|
||||||
"""Validate that the dataset class inherits from BaseDataset."""
|
"""Validate that the dataset class inherits from BaseDataset."""
|
||||||
if not issubclass(dataset_cls, BaseDataset):
|
if not issubclass(dataset_cls, BaseDataset):
|
||||||
raise TypeError(f"{dataset_cls.__name__} must inherit from BaseDataset")
|
raise TypeError(f"{dataset_cls.__name__} must inherit from BaseDataset")
|
||||||
@@ -175,7 +164,6 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
window_size: int,
|
window_size: int,
|
||||||
stride: Optional[int] = None,
|
stride: Optional[int] = None,
|
||||||
storage_type: Optional[str] = None,
|
storage_type: Optional[str] = None,
|
||||||
tokenizer=None,
|
|
||||||
) -> "BaseDataset":
|
) -> "BaseDataset":
|
||||||
"""Create and load a dataset in one step.
|
"""Create and load a dataset in one step.
|
||||||
|
|
||||||
@@ -184,8 +172,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
load_path: Path to the data file
|
load_path: Path to the data file
|
||||||
window_size: Window size for data sampling
|
window_size: Window size for data sampling
|
||||||
stride: Stride between consecutive samples (default: same as window_size)
|
stride: Stride between consecutive samples (default: same as window_size)
|
||||||
storage_type: Storage type ("h5", "json") or None for auto-detection
|
storage_type: Storage type ("h5", "bin") or None for auto-detection
|
||||||
tokenizer: Callable str -> List[int] for raw text JSON tokenization
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Loaded dataset instance
|
Loaded dataset instance
|
||||||
@@ -194,7 +181,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
stride = window_size
|
stride = window_size
|
||||||
|
|
||||||
dataset = cls.create(train_type, window_size, stride)
|
dataset = cls.create(train_type, window_size, stride)
|
||||||
dataset.load(load_path, storage_type=storage_type, tokenizer=tokenizer)
|
dataset.load(load_path, storage_type=storage_type)
|
||||||
|
|
||||||
return dataset
|
return dataset
|
||||||
|
|
||||||
@@ -306,9 +293,11 @@ class GRPODataset(BaseDataset):
|
|||||||
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, end_idx = self.get_index(index)
|
||||||
|
|
||||||
prompts = self._fetch_data(begin_idx, end_idx, "prompts")
|
prompts = self._fetch_data(begin_idx, end_idx, "prompts").to(dtype=torch.long)
|
||||||
responses = self._fetch_data(begin_idx, end_idx, "responses")
|
responses = self._fetch_data(begin_idx, end_idx, "responses").to(
|
||||||
masks = self._fetch_data(begin_idx, end_idx, "masks")
|
dtype=torch.long
|
||||||
|
)
|
||||||
|
masks = self._fetch_data(begin_idx, end_idx, "masks").to(dtype=torch.bool)
|
||||||
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
||||||
|
|
||||||
return {
|
return {
|
||||||
|
|||||||
@@ -43,6 +43,7 @@ class ResumableDistributedSampler(Sampler[int]):
|
|||||||
offset = 0 if drop_last else self.num_replicas - 1
|
offset = 0 if drop_last else self.num_replicas - 1
|
||||||
self.num_samples_per_replica = (self.num_samples + offset) // self.num_replicas
|
self.num_samples_per_replica = (self.num_samples + offset) // self.num_replicas
|
||||||
self.total_size = self.num_samples_per_replica * self.num_replicas
|
self.total_size = self.num_samples_per_replica * self.num_replicas
|
||||||
|
self.iter = self.iter % self.num_samples_per_replica
|
||||||
|
|
||||||
self._indices = None
|
self._indices = None
|
||||||
|
|
||||||
@@ -74,5 +75,10 @@ class ResumableDistributedSampler(Sampler[int]):
|
|||||||
self.epoch += 1
|
self.epoch += 1
|
||||||
self._indices = None
|
self._indices = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def _remaining(self):
|
||||||
|
remaining = self.num_samples_per_replica - self.iter
|
||||||
|
return max(remaining, 0)
|
||||||
|
|
||||||
def __len__(self):
|
def __len__(self):
|
||||||
return self.num_samples_per_replica
|
return self._remaining
|
||||||
|
|||||||
+140
-191
@@ -1,7 +1,20 @@
|
|||||||
"""Storage backends for different data formats.
|
"""Storage backends for different data formats.
|
||||||
|
|
||||||
Each storage handles format-specific loading (HDF5, JSON, etc.) and provides
|
Layers:
|
||||||
a uniform interface for data access and length observation via fetchers.
|
- I/O layer: save_* / load_* functions, read/write raw files (HDF5/bin)
|
||||||
|
return Dict[str, List[Tensor]] — format-specific, no state
|
||||||
|
- Store (ABC): central abstraction, normalizes multi-segment into
|
||||||
|
Dict[str, List[Tensor]] per key via _normalize(),
|
||||||
|
fetch() uses bisect across segments — no forced concat
|
||||||
|
- Dataset layer: BaseDataset owns a Store, only calls store.fetch(begin, end, key)
|
||||||
|
|
||||||
|
Key properties:
|
||||||
|
- Multi-segment: segments kept as-is, no forced concatenation — safe for
|
||||||
|
datasets larger than RAM
|
||||||
|
- Explicit length: _length = min(total elements across keys), set at load,
|
||||||
|
__len__ returns O(1)
|
||||||
|
- Zero-copy mmap: MmapStore wraps np.memmap(mode="r"), all DataLoader
|
||||||
|
workers share OS page-cache pages
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import bisect
|
import bisect
|
||||||
@@ -9,9 +22,10 @@ import json
|
|||||||
import os
|
import os
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Callable, Dict, List, Optional, Union
|
from typing import Dict, List, Union
|
||||||
|
|
||||||
import h5py
|
import h5py
|
||||||
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
@@ -54,54 +68,30 @@ def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
|||||||
return tensor_group
|
return tensor_group
|
||||||
|
|
||||||
|
|
||||||
def save_json(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
def save_bin(file_path: str, tensor_group: Dict[str, List[Tensor]]):
|
||||||
os.makedirs(file_path, exist_ok=True)
|
os.makedirs(file_path, exist_ok=True)
|
||||||
full_file_path = os.path.join(file_path, f"{file_name}.json")
|
meta = {}
|
||||||
json_data = {}
|
|
||||||
for key, tensors in tensor_group.items():
|
for key, tensors in tensor_group.items():
|
||||||
json_data[key] = [tensor.tolist() for tensor in tensors]
|
cat = torch.cat(tensors, dim=0)
|
||||||
with open(full_file_path, "w", encoding="utf-8") as f:
|
meta[key] = {"shape": list(cat.shape), "dtype": str(cat.dtype).split(".")[-1]}
|
||||||
json.dump(json_data, f, ensure_ascii=False)
|
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
||||||
|
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
||||||
|
json.dump(meta, f)
|
||||||
|
|
||||||
|
|
||||||
def load_json(
|
def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
||||||
file_path: str,
|
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||||
share_memory: bool = True,
|
meta = json.load(f)
|
||||||
tokenizer: Optional[Callable[[str], List[int]]] = None,
|
segments: Dict[str, List[Tensor]] = {}
|
||||||
) -> Dict[str, List[Tensor]]:
|
for key, info in meta.items():
|
||||||
"""Load tensor data from JSON files.
|
arr = np.memmap(
|
||||||
|
os.path.join(file_path, f"{key}.bin"),
|
||||||
Supports two modes:
|
dtype=info["dtype"],
|
||||||
- Pre-tokenized: JSON values are List[List[int]] (token IDs), loaded as-is.
|
mode="r+",
|
||||||
- Raw text: JSON values are List[str], tokenized via ``tokenizer`` callable
|
shape=tuple(info["shape"]),
|
||||||
at load time. A ``tokenizer`` receives a str and returns List[int].
|
)
|
||||||
|
segments[key] = [torch.from_numpy(arr)]
|
||||||
Non-data JSON files (e.g. config.json) with scalar/object values are
|
return segments
|
||||||
silently skipped.
|
|
||||||
"""
|
|
||||||
tensor_group: Dict[str, List[Tensor]] = {}
|
|
||||||
root_path = Path(file_path)
|
|
||||||
json_files = list(root_path.rglob("*.json")) + list(root_path.rglob("*.jsonl"))
|
|
||||||
for json_file in json_files:
|
|
||||||
with open(json_file, "r", encoding="utf-8") as f:
|
|
||||||
data = json.load(f)
|
|
||||||
if not isinstance(data, dict):
|
|
||||||
continue
|
|
||||||
for key, sequences in data.items():
|
|
||||||
if not isinstance(sequences, list):
|
|
||||||
continue
|
|
||||||
tensors = []
|
|
||||||
for seq in sequences:
|
|
||||||
if tokenizer is not None and isinstance(seq, str):
|
|
||||||
seq = tokenizer(seq)
|
|
||||||
tensor = torch.tensor(seq, dtype=torch.long)
|
|
||||||
if share_memory:
|
|
||||||
tensor = tensor.share_memory_()
|
|
||||||
tensors.append(tensor)
|
|
||||||
if tensor_group.get(key) is None:
|
|
||||||
tensor_group[key] = []
|
|
||||||
tensor_group[key].extend(tensors)
|
|
||||||
return tensor_group
|
|
||||||
|
|
||||||
|
|
||||||
def detect_format(load_path: str) -> str:
|
def detect_format(load_path: str) -> str:
|
||||||
@@ -111,7 +101,7 @@ def detect_format(load_path: str) -> str:
|
|||||||
load_path: Directory or file path
|
load_path: Directory or file path
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Format string ("h5" or "json")
|
Format string ("h5" or "bin")
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
FileNotFoundError: If no supported data files are found
|
FileNotFoundError: If no supported data files are found
|
||||||
@@ -121,181 +111,140 @@ def detect_format(load_path: str) -> str:
|
|||||||
suffix = root.suffix.lower()
|
suffix = root.suffix.lower()
|
||||||
if suffix in (".h5", ".hdf5"):
|
if suffix in (".h5", ".hdf5"):
|
||||||
return "h5"
|
return "h5"
|
||||||
if suffix in (".json", ".jsonl"):
|
|
||||||
return "json"
|
|
||||||
raise ValueError(f"Unsupported file format: {suffix}")
|
raise ValueError(f"Unsupported file format: {suffix}")
|
||||||
|
|
||||||
h5_files = list(root.rglob("*.h5")) + list(root.rglob("*.hdf5"))
|
h5_files = list(root.rglob("*.h5")) + list(root.rglob("*.hdf5"))
|
||||||
if h5_files:
|
if h5_files:
|
||||||
return "h5"
|
return "h5"
|
||||||
json_files = list(root.rglob("*.json")) + list(root.rglob("*.jsonl"))
|
bin_files = list(root.rglob("*.bin"))
|
||||||
if json_files:
|
if bin_files and (root / "meta.json").exists():
|
||||||
return "json"
|
return "bin"
|
||||||
raise FileNotFoundError(f"No supported data files found at {load_path}")
|
raise FileNotFoundError(f"No supported data files found at {load_path}")
|
||||||
|
|
||||||
|
|
||||||
class BaseSegmentFetcher:
|
class Store(ABC):
|
||||||
"""Fetches data segments across multiple tensor segments.
|
"""String keys -> segmented tensors with ``fetch(begin, end, keys)``.
|
||||||
|
|
||||||
Maintains cumulative lengths for efficient range queries across
|
Each key maps to one or more tensor segments (no forced concatenation).
|
||||||
multiple discontinuous segments.
|
``len(store)`` returns ``self._length`` (explicit, O(1)), the minimum
|
||||||
"""
|
total element count across all keys.
|
||||||
|
|
||||||
def __init__(self, segments: List[Tensor]):
|
Subclasses fill ``self._data`` and ``self._cum`` during ``load()``
|
||||||
self.segments = segments
|
via ``_normalize()``.
|
||||||
self.cum_lengths = []
|
|
||||||
|
|
||||||
total = 0
|
|
||||||
for seg in segments:
|
|
||||||
total += torch.numel(seg)
|
|
||||||
self.cum_lengths.append(total)
|
|
||||||
|
|
||||||
self.total_length = total
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
return self.total_length
|
|
||||||
|
|
||||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
|
||||||
"""Fetch data in the range [begin_idx, end_idx)."""
|
|
||||||
if not (
|
|
||||||
0 <= begin_idx < self.total_length and 0 <= end_idx <= self.total_length
|
|
||||||
):
|
|
||||||
raise ValueError("begin_idx or end_idx out of bounds")
|
|
||||||
if begin_idx >= end_idx:
|
|
||||||
return torch.tensor([], dtype=torch.long)
|
|
||||||
|
|
||||||
seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx)
|
|
||||||
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx)
|
|
||||||
|
|
||||||
result_segments = []
|
|
||||||
|
|
||||||
for i in range(seg_start_idx, seg_end_idx + 1):
|
|
||||||
prev_cum = self.cum_lengths[i - 1] if i > 0 else 0
|
|
||||||
start = max(begin_idx - prev_cum, 0)
|
|
||||||
end = min(end_idx - prev_cum, len(self.segments[i]))
|
|
||||||
result_segments.append(self.segments[i][start:end])
|
|
||||||
|
|
||||||
return torch.cat(result_segments, dim=0)
|
|
||||||
|
|
||||||
|
|
||||||
class MultiSegmentFetcher:
|
|
||||||
"""Manages multiple segment fetchers for different data keys."""
|
|
||||||
|
|
||||||
def __init__(self, multi_segments: Dict):
|
|
||||||
self.multi_keys = list(multi_segments.keys())
|
|
||||||
self.multi_fetchers = {
|
|
||||||
key: BaseSegmentFetcher(segments)
|
|
||||||
for key, segments in multi_segments.items()
|
|
||||||
}
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
"""Returns the minimum length across all fetchers."""
|
|
||||||
if not self.multi_fetchers:
|
|
||||||
return 0
|
|
||||||
len_list = [len(seg) for seg in self.multi_fetchers.values()]
|
|
||||||
return min(len_list)
|
|
||||||
|
|
||||||
def key_fetch(
|
|
||||||
self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]
|
|
||||||
) -> Dict:
|
|
||||||
"""Fetch data for specific keys."""
|
|
||||||
fetch_dict = {}
|
|
||||||
keys = [keys] if isinstance(keys, str) else keys
|
|
||||||
|
|
||||||
for key in keys:
|
|
||||||
fetcher = self.multi_fetchers[key]
|
|
||||||
fetch_tensor = fetcher.fetch_data(begin_idx, end_idx)
|
|
||||||
fetch_dict[key] = fetch_tensor
|
|
||||||
|
|
||||||
return fetch_dict if len(keys) > 1 else fetch_dict[keys[0]]
|
|
||||||
|
|
||||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Dict:
|
|
||||||
"""Fetch all keys."""
|
|
||||||
return self.key_fetch(begin_idx, end_idx, self.multi_keys)
|
|
||||||
|
|
||||||
|
|
||||||
class BaseStorage(ABC):
|
|
||||||
"""Abstract storage backend for loading and dispatching data.
|
|
||||||
|
|
||||||
Storage encapsulates format-specific loading and provides a uniform
|
|
||||||
interface for data access and length observation. Subclasses handle
|
|
||||||
different data formats (HDF5, JSON, etc.) while exposing the same
|
|
||||||
fetch interface.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self._fetcher: Optional[MultiSegmentFetcher] = None
|
self._data: Dict[str, List[Tensor]] = {}
|
||||||
|
self._cum: Dict[str, List[int]] = {}
|
||||||
|
self._length: int = 0
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def load(self, load_path: str, tokenizer=None) -> None:
|
def load(self, path: str) -> None:
|
||||||
"""Load data from the given path into internal fetcher."""
|
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
"""Total number of raw elements (tokens) in storage."""
|
|
||||||
if self._fetcher is None:
|
|
||||||
return 0
|
|
||||||
return len(self._fetcher)
|
|
||||||
|
|
||||||
def fetch(self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]):
|
|
||||||
"""Fetch data for the given keys and index range.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
begin_idx: Starting index (inclusive)
|
|
||||||
end_idx: Ending index (exclusive)
|
|
||||||
keys: Single key or list of keys to fetch
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tensor if single key, Dict[str, Tensor] if multiple keys
|
|
||||||
"""
|
|
||||||
if self._fetcher is None:
|
|
||||||
raise RuntimeError("Storage not loaded")
|
|
||||||
return self._fetcher.key_fetch(begin_idx, end_idx, keys)
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def keys(self) -> List[str]:
|
def keys(self) -> List[str]:
|
||||||
"""Return the data keys available in this storage."""
|
return list(self._data.keys())
|
||||||
if self._fetcher is None:
|
|
||||||
return []
|
def __len__(self) -> int:
|
||||||
return self._fetcher.multi_keys
|
return self._length
|
||||||
|
|
||||||
|
def fetch(
|
||||||
|
self,
|
||||||
|
begin: int,
|
||||||
|
end: int,
|
||||||
|
keys: Union[str, List[str]],
|
||||||
|
):
|
||||||
|
if not self._data:
|
||||||
|
raise RuntimeError("Store not loaded")
|
||||||
|
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
||||||
|
raise ValueError(
|
||||||
|
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
||||||
|
)
|
||||||
|
if isinstance(keys, str):
|
||||||
|
return self._fetch_key(keys, begin, end)
|
||||||
|
return {k: self._fetch_key(k, begin, end) for k in keys}
|
||||||
|
|
||||||
|
def _fetch_key(self, key: str, begin: int, end: int) -> Tensor:
|
||||||
|
"""Fetch slice [begin, end) across potentially multiple segments."""
|
||||||
|
segments = self._data[key]
|
||||||
|
cum = self._cum[key]
|
||||||
|
seg_start = bisect.bisect_right(cum, begin)
|
||||||
|
seg_end = bisect.bisect_left(cum, end)
|
||||||
|
|
||||||
|
results = []
|
||||||
|
for i in range(seg_start, seg_end + 1):
|
||||||
|
prev = cum[i - 1] if i > 0 else 0
|
||||||
|
s = max(begin - prev, 0)
|
||||||
|
e = min(end - prev, segments[i].shape[0])
|
||||||
|
results.append(segments[i][s:e])
|
||||||
|
|
||||||
|
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
|
||||||
|
|
||||||
|
def _normalize(self, raw: Dict[str, List[Tensor]]):
|
||||||
|
"""Register segments and pre-compute cumulative lengths.
|
||||||
|
|
||||||
|
Does NOT concatenate — segments are kept as-is to avoid OOM on
|
||||||
|
large datasets. Sets ``self._length`` to the minimum total
|
||||||
|
element count across all keys.
|
||||||
|
"""
|
||||||
|
for key, tensors in raw.items():
|
||||||
|
self._data[key] = tensors
|
||||||
|
cum = []
|
||||||
|
total = 0
|
||||||
|
for t in tensors:
|
||||||
|
total += t.shape[0]
|
||||||
|
cum.append(total)
|
||||||
|
self._cum[key] = cum
|
||||||
|
self._length = (
|
||||||
|
min((cum[-1] if cum else 0) for cum in self._cum.values())
|
||||||
|
if self._cum
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class StorageFactory(BaseFactory["BaseStorage"]):
|
class StoreFactory(BaseFactory["Store"]):
|
||||||
"""Factory for creating storage backends by type name.
|
"""Factory for creating Store instances by type name.
|
||||||
|
|
||||||
Example:
|
Example::
|
||||||
@StorageFactory.register("custom")
|
|
||||||
class CustomStorage(BaseStorage):
|
@StoreFactory.register("custom")
|
||||||
|
class CustomStore(Store):
|
||||||
...
|
...
|
||||||
|
|
||||||
storage = StorageFactory.create("custom")
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _validate_component(cls, storage_cls: type) -> None:
|
def _validate_component(cls, store_cls: type):
|
||||||
if not issubclass(storage_cls, BaseStorage):
|
if not issubclass(store_cls, Store):
|
||||||
raise TypeError(f"{storage_cls.__name__} must inherit from BaseStorage")
|
raise TypeError(f"{store_cls.__name__} must inherit from Store")
|
||||||
|
|
||||||
|
|
||||||
@StorageFactory.register("h5")
|
@StoreFactory.register("h5")
|
||||||
class H5Storage(BaseStorage):
|
class H5Store(Store):
|
||||||
"""HDF5-based storage backend (pre-tokenized data)."""
|
"""HDF5-based storage backend (pre-tokenized data)."""
|
||||||
|
|
||||||
def load(self, load_path: str, tokenizer=None) -> None:
|
def load(self, path: str):
|
||||||
segments = load_h5(load_path)
|
self._normalize(load_h5(path))
|
||||||
self._fetcher = MultiSegmentFetcher(segments)
|
|
||||||
|
|
||||||
|
|
||||||
@StorageFactory.register("json")
|
@StoreFactory.register("bin")
|
||||||
class JSONStorage(BaseStorage):
|
class MmapStore(Store):
|
||||||
"""JSON-based storage backend.
|
"""Memory-mapped binary storage backend.
|
||||||
|
|
||||||
Supports two modes:
|
Each key is a single .bin file backed by ``np.memmap(mode="r")``.
|
||||||
- Pre-tokenized: JSON values are List[List[int]], loaded as-is.
|
No per-process memory duplication — all DataLoader workers share the
|
||||||
- Raw text: JSON values are List[str], tokenized via ``tokenizer``
|
same OS page-cache pages.
|
||||||
callable (str -> List[int]) at load time.
|
|
||||||
|
Format on disk::
|
||||||
|
|
||||||
|
data_root/
|
||||||
|
meta.json # {key: {shape, dtype}, ...}
|
||||||
|
<key>.bin # raw numpy array, one per key
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def load(self, load_path: str, tokenizer=None) -> None:
|
def load(self, path: str):
|
||||||
segments = load_json(load_path, tokenizer=tokenizer)
|
self._mmap_refs = []
|
||||||
self._fetcher = MultiSegmentFetcher(segments)
|
raw = load_bin(path)
|
||||||
|
self._normalize(raw)
|
||||||
|
for tensors in self._data.values():
|
||||||
|
self._mmap_refs.extend(tensors)
|
||||||
|
|||||||
+2
-2
@@ -23,7 +23,7 @@ class Registry:
|
|||||||
component_cls: Type,
|
component_cls: Type,
|
||||||
category: Optional[str] = None,
|
category: Optional[str] = None,
|
||||||
priority: int = 0,
|
priority: int = 0,
|
||||||
) -> None:
|
):
|
||||||
"""Register a component class with optional category and priority."""
|
"""Register a component class with optional category and priority."""
|
||||||
if name in self._entries:
|
if name in self._entries:
|
||||||
raise ValueError(f"Component '{name}' is already registered")
|
raise ValueError(f"Component '{name}' is already registered")
|
||||||
@@ -158,7 +158,7 @@ class BaseFactory(ABC, Generic[T]):
|
|||||||
return component_cls(*args, **kwargs)
|
return component_cls(*args, **kwargs)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _validate_component(cls, component_cls: Type[T]) -> None:
|
def _validate_component(cls, component_cls: Type[T]):
|
||||||
"""Validate that the component class is valid for this factory.
|
"""Validate that the component class is valid for this factory.
|
||||||
|
|
||||||
Override this method in subclasses to add custom validation.
|
Override this method in subclasses to add custom validation.
|
||||||
|
|||||||
@@ -2,24 +2,26 @@
|
|||||||
|
|
||||||
Layers:
|
Layers:
|
||||||
- core/: Core inference loop (cache, executor, scheduler, task)
|
- core/: Core inference loop (cache, executor, scheduler, task)
|
||||||
- api/: HTTP protocol handlers (OpenAI, Anthropic)
|
- api/: HTTP orchestration (ProtocolHandler, server)
|
||||||
|
- protocols/: Response builders (OpenAI, Anthropic)
|
||||||
|
- transport/: SSE transport utilities
|
||||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from astrai.inference.api import (
|
from astrai.inference.api import (
|
||||||
AnthropicHandler,
|
|
||||||
AnthropicMessage,
|
AnthropicMessage,
|
||||||
ChatCompletionRequest,
|
ChatCompletionRequest,
|
||||||
ChatMessage,
|
ChatMessage,
|
||||||
|
GenContext,
|
||||||
MessagesRequest,
|
MessagesRequest,
|
||||||
OpenAIHandler,
|
|
||||||
ProtocolHandler,
|
ProtocolHandler,
|
||||||
StopChecker,
|
StopChecker,
|
||||||
StreamContext,
|
|
||||||
app,
|
app,
|
||||||
run_server,
|
run_server,
|
||||||
)
|
)
|
||||||
|
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||||
|
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||||
from astrai.inference.core import (
|
from astrai.inference.core import (
|
||||||
STOP,
|
STOP,
|
||||||
Allocator,
|
Allocator,
|
||||||
@@ -36,10 +38,7 @@ from astrai.inference.core import (
|
|||||||
TaskTable,
|
TaskTable,
|
||||||
page_hash,
|
page_hash,
|
||||||
)
|
)
|
||||||
from astrai.inference.engine import (
|
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
||||||
GenerationRequest,
|
|
||||||
InferenceEngine,
|
|
||||||
)
|
|
||||||
from astrai.inference.sample import (
|
from astrai.inference.sample import (
|
||||||
BaseSamplingStrategy,
|
BaseSamplingStrategy,
|
||||||
SamplingPipeline,
|
SamplingPipeline,
|
||||||
@@ -50,17 +49,14 @@ from astrai.inference.sample import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Engine / Requests
|
|
||||||
"InferenceEngine",
|
"InferenceEngine",
|
||||||
"GenerationRequest",
|
"GenerationRequest",
|
||||||
# Core scheduler
|
|
||||||
"InferenceScheduler",
|
"InferenceScheduler",
|
||||||
"Executor",
|
"Executor",
|
||||||
"STOP",
|
"STOP",
|
||||||
"Task",
|
"Task",
|
||||||
"TaskManager",
|
"TaskManager",
|
||||||
"TaskStatus",
|
"TaskStatus",
|
||||||
# Core cache
|
|
||||||
"Allocator",
|
"Allocator",
|
||||||
"KVCache",
|
"KVCache",
|
||||||
"KvcacheView",
|
"KvcacheView",
|
||||||
@@ -69,20 +65,17 @@ __all__ = [
|
|||||||
"Storage",
|
"Storage",
|
||||||
"TaskTable",
|
"TaskTable",
|
||||||
"page_hash",
|
"page_hash",
|
||||||
# Sampling (Strategy pattern)
|
|
||||||
"sample",
|
"sample",
|
||||||
"BaseSamplingStrategy",
|
"BaseSamplingStrategy",
|
||||||
"TemperatureStrategy",
|
"TemperatureStrategy",
|
||||||
"TopKStrategy",
|
"TopKStrategy",
|
||||||
"TopPStrategy",
|
"TopPStrategy",
|
||||||
"SamplingPipeline",
|
"SamplingPipeline",
|
||||||
# Protocol
|
|
||||||
"ProtocolHandler",
|
"ProtocolHandler",
|
||||||
"StopChecker",
|
"StopChecker",
|
||||||
"StreamContext",
|
"GenContext",
|
||||||
"AnthropicHandler",
|
"OpenAIResponseBuilder",
|
||||||
"OpenAIHandler",
|
"AnthropicResponseBuilder",
|
||||||
# Server
|
|
||||||
"ChatMessage",
|
"ChatMessage",
|
||||||
"ChatCompletionRequest",
|
"ChatCompletionRequest",
|
||||||
"AnthropicMessage",
|
"AnthropicMessage",
|
||||||
|
|||||||
@@ -1,12 +1,6 @@
|
|||||||
"""Inference API: protocol handlers and FastAPI server."""
|
"""Inference API: protocol handler, stop checker, and FastAPI server."""
|
||||||
|
|
||||||
from astrai.inference.api.protocol import (
|
from astrai.inference.api.protocol import GenContext, ProtocolHandler, StopChecker
|
||||||
AnthropicHandler,
|
|
||||||
OpenAIHandler,
|
|
||||||
ProtocolHandler,
|
|
||||||
StopChecker,
|
|
||||||
StreamContext,
|
|
||||||
)
|
|
||||||
from astrai.inference.api.server import (
|
from astrai.inference.api.server import (
|
||||||
AnthropicMessage,
|
AnthropicMessage,
|
||||||
ChatCompletionRequest,
|
ChatCompletionRequest,
|
||||||
@@ -17,11 +11,9 @@ from astrai.inference.api.server import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AnthropicHandler",
|
|
||||||
"OpenAIHandler",
|
|
||||||
"ProtocolHandler",
|
"ProtocolHandler",
|
||||||
"StopChecker",
|
"StopChecker",
|
||||||
"StreamContext",
|
"GenContext",
|
||||||
"AnthropicMessage",
|
"AnthropicMessage",
|
||||||
"ChatCompletionRequest",
|
"ChatCompletionRequest",
|
||||||
"ChatMessage",
|
"ChatMessage",
|
||||||
|
|||||||
@@ -0,0 +1,140 @@
|
|||||||
|
"""Anthropic message completion response builder."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from typing import Any, Dict, List, Tuple, Union
|
||||||
|
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from astrai.inference.api.protocol import (
|
||||||
|
GenContext,
|
||||||
|
ResponseBuilder,
|
||||||
|
StopInfo,
|
||||||
|
sse_event,
|
||||||
|
)
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_text(content: Union[str, List[Dict[str, Any]]]) -> str:
|
||||||
|
if isinstance(content, str):
|
||||||
|
return content
|
||||||
|
if isinstance(content, list):
|
||||||
|
for block in content:
|
||||||
|
if isinstance(block, dict) and block.get("type") == "text":
|
||||||
|
return block.get("text", "")
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
class AnthropicResponseBuilder(ResponseBuilder):
|
||||||
|
def prepare(
|
||||||
|
self, request: BaseModel, engine: InferenceEngine
|
||||||
|
) -> Tuple[str, GenContext, List[str]]:
|
||||||
|
messages: List[Dict[str, str]] = []
|
||||||
|
system = getattr(request, "system", None)
|
||||||
|
if system:
|
||||||
|
messages.append({"role": "system", "content": system})
|
||||||
|
for m in request.messages:
|
||||||
|
text = _extract_text(m.content)
|
||||||
|
if text:
|
||||||
|
messages.append({"role": m.role, "content": text})
|
||||||
|
prompt = engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||||
|
ctx = GenContext(
|
||||||
|
resp_id=f"msg_{uuid.uuid4().hex[:24]}",
|
||||||
|
created=0,
|
||||||
|
model=request.model,
|
||||||
|
prompt_tokens=0,
|
||||||
|
)
|
||||||
|
stop_sequences = getattr(request, "stop_sequences", None) or []
|
||||||
|
return prompt, ctx, stop_sequences
|
||||||
|
|
||||||
|
def format_stream_start(self, ctx: GenContext) -> List[str]:
|
||||||
|
return [
|
||||||
|
sse_event(
|
||||||
|
{
|
||||||
|
"type": "message_start",
|
||||||
|
"message": {
|
||||||
|
"id": ctx.resp_id,
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": ctx.model,
|
||||||
|
"content": [],
|
||||||
|
"usage": {"input_tokens": ctx.prompt_tokens},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
event="message_start",
|
||||||
|
),
|
||||||
|
sse_event(
|
||||||
|
{
|
||||||
|
"type": "content_block_start",
|
||||||
|
"index": 0,
|
||||||
|
"content_block": {"type": "text", "text": ""},
|
||||||
|
},
|
||||||
|
event="content_block_start",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
def format_chunk(self, token: str) -> str:
|
||||||
|
return sse_event(
|
||||||
|
{
|
||||||
|
"type": "content_block_delta",
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"type": "text_delta", "text": token},
|
||||||
|
},
|
||||||
|
event="content_block_delta",
|
||||||
|
)
|
||||||
|
|
||||||
|
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||||
|
events: List[str] = []
|
||||||
|
if stop.matched:
|
||||||
|
trimmed = stop.body[: stop.body.rfind(stop.matched)]
|
||||||
|
unyielded = trimmed[len(stop.yielded) :]
|
||||||
|
if unyielded:
|
||||||
|
events.append(
|
||||||
|
sse_event(
|
||||||
|
{
|
||||||
|
"type": "content_block_delta",
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"type": "text_delta", "text": unyielded},
|
||||||
|
},
|
||||||
|
event="content_block_delta",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
events.append(
|
||||||
|
sse_event(
|
||||||
|
{"type": "content_block_stop", "index": 0},
|
||||||
|
event="content_block_stop",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
events.append(
|
||||||
|
sse_event(
|
||||||
|
{
|
||||||
|
"type": "message_delta",
|
||||||
|
"delta": {
|
||||||
|
"stop_reason": "stop_sequence" if stop.matched else "end_turn",
|
||||||
|
"stop_sequence": stop.matched,
|
||||||
|
},
|
||||||
|
"usage": {"output_tokens": ctx.completion_tokens},
|
||||||
|
},
|
||||||
|
event="message_delta",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
events.append(sse_event({"type": "message_stop"}, event="message_stop"))
|
||||||
|
return events
|
||||||
|
|
||||||
|
def format_response(
|
||||||
|
self, ctx: GenContext, content: str, stop: StopInfo
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
if stop.matched:
|
||||||
|
content = content[: content.rfind(stop.matched)]
|
||||||
|
return {
|
||||||
|
"id": ctx.resp_id,
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": ctx.model,
|
||||||
|
"content": [{"type": "text", "text": content}],
|
||||||
|
"stop_reason": "stop_sequence" if stop.matched else "end_turn",
|
||||||
|
"stop_sequence": stop.matched,
|
||||||
|
"usage": {
|
||||||
|
"input_tokens": ctx.prompt_tokens,
|
||||||
|
"output_tokens": ctx.completion_tokens,
|
||||||
|
},
|
||||||
|
}
|
||||||
@@ -0,0 +1,111 @@
|
|||||||
|
"""OpenAI chat completion response builder."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from typing import Any, Dict, List, Tuple
|
||||||
|
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from astrai.inference.api.protocol import (
|
||||||
|
GenContext,
|
||||||
|
ResponseBuilder,
|
||||||
|
StopInfo,
|
||||||
|
sse_event,
|
||||||
|
)
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
|
||||||
|
class OpenAIResponseBuilder(ResponseBuilder):
|
||||||
|
def prepare(
|
||||||
|
self, request: BaseModel, engine: InferenceEngine
|
||||||
|
) -> Tuple[str, GenContext, List[str]]:
|
||||||
|
messages = [{"role": m.role, "content": m.content} for m in request.messages]
|
||||||
|
prompt = engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||||
|
|
||||||
|
self._resp_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||||
|
self._model = request.model
|
||||||
|
|
||||||
|
ctx = GenContext(
|
||||||
|
resp_id=self._resp_id,
|
||||||
|
created=0,
|
||||||
|
model=self._model,
|
||||||
|
prompt_tokens=0,
|
||||||
|
)
|
||||||
|
stop = request.stop
|
||||||
|
stop_sequences = (
|
||||||
|
[] if stop is None else [stop] if isinstance(stop, str) else stop
|
||||||
|
)
|
||||||
|
return prompt, ctx, stop_sequences
|
||||||
|
|
||||||
|
def format_stream_start(self, ctx: GenContext) -> List[str]:
|
||||||
|
return [
|
||||||
|
sse_event(
|
||||||
|
{
|
||||||
|
"id": self._resp_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": ctx.created,
|
||||||
|
"model": self._model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"role": "assistant"},
|
||||||
|
"finish_reason": None,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
def format_chunk(self, token: str) -> str:
|
||||||
|
return sse_event(
|
||||||
|
{
|
||||||
|
"id": self._resp_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": 0,
|
||||||
|
"model": self._model,
|
||||||
|
"choices": [
|
||||||
|
{"index": 0, "delta": {"content": token}, "finish_reason": None}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||||
|
return [
|
||||||
|
sse_event(
|
||||||
|
{
|
||||||
|
"id": self._resp_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": ctx.created,
|
||||||
|
"model": self._model,
|
||||||
|
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||||
|
}
|
||||||
|
),
|
||||||
|
sse_event(
|
||||||
|
{
|
||||||
|
"prompt_tokens": ctx.prompt_tokens,
|
||||||
|
"completion_tokens": ctx.completion_tokens,
|
||||||
|
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||||
|
}
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
def format_response(
|
||||||
|
self, ctx: GenContext, content: str, stop: StopInfo
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"id": self._resp_id,
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": ctx.created,
|
||||||
|
"model": self._model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"message": {"role": "assistant", "content": content},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": ctx.prompt_tokens,
|
||||||
|
"completion_tokens": ctx.completion_tokens,
|
||||||
|
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||||
|
},
|
||||||
|
}
|
||||||
@@ -1,15 +1,13 @@
|
|||||||
"""Protocol handlers for OpenAI and Anthropic chat completion APIs.
|
"""Orchestration layer: ProtocolHandler, StopChecker, GenContext, StopInfo, ResponseBuilder, SSE utils.
|
||||||
|
|
||||||
Template Method + Builder patterns eliminate the 45% code duplication between
|
ProtocolHandler orchestrates the async generation loop and delegates
|
||||||
stream/non-stream branches and across protocol adapters.
|
protocol-specific formatting to a ResponseBuilder.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import time
|
|
||||||
import uuid
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Dict, List, Optional, Union
|
from typing import Any, AsyncGenerator, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
from fastapi.responses import StreamingResponse
|
from fastapi.responses import StreamingResponse
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
@@ -17,7 +15,7 @@ from pydantic import BaseModel
|
|||||||
from astrai.inference.engine import InferenceEngine
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
|
||||||
def _sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
def sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
||||||
lines: List[str] = []
|
lines: List[str] = []
|
||||||
if event:
|
if event:
|
||||||
lines.append(f"event: {event}")
|
lines.append(f"event: {event}")
|
||||||
@@ -26,22 +24,28 @@ def _sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
|||||||
return "\n".join(lines)
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
def _sse_done() -> str:
|
def sse_done() -> str:
|
||||||
return "data: [DONE]\n\n"
|
return "data: [DONE]\n\n"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class StreamContext:
|
class GenContext:
|
||||||
"""Shared state across the streaming generation lifecycle."""
|
"""Per-generation metadata passed to builder format methods."""
|
||||||
|
|
||||||
resp_id: str
|
resp_id: str
|
||||||
created: int
|
created: int
|
||||||
model: str
|
model: str
|
||||||
prompt_tokens: int
|
prompt_tokens: int
|
||||||
completion_tokens: int = 0
|
completion_tokens: int = 0
|
||||||
accumulated: str = ""
|
|
||||||
stop_matched: Optional[str] = None
|
|
||||||
last_yield_trimmed: str = ""
|
@dataclass
|
||||||
|
class StopInfo:
|
||||||
|
"""Stop-check result passed to format_stream_end / format_response."""
|
||||||
|
|
||||||
|
matched: Optional[str] = None
|
||||||
|
body: str = ""
|
||||||
|
yielded: str = ""
|
||||||
|
|
||||||
|
|
||||||
class StopChecker:
|
class StopChecker:
|
||||||
@@ -56,95 +60,60 @@ class StopChecker:
|
|||||||
return seq
|
return seq
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def trim(self, text: str, matched: str) -> str:
|
|
||||||
idx = text.rfind(matched)
|
|
||||||
return text[:idx] if idx != -1 else text
|
|
||||||
|
|
||||||
@property
|
class ResponseBuilder(ABC):
|
||||||
def has_sequences(self) -> bool:
|
"""Interface for protocol-specific response formatting.
|
||||||
return len(self._sequences) > 0
|
|
||||||
|
|
||||||
|
A new protocol requires one concrete builder implementing 6 methods.
|
||||||
class ProtocolHandler(ABC):
|
|
||||||
"""Template-method base for API protocol handlers.
|
|
||||||
|
|
||||||
Subclasses implement format hooks; the base class orchestrates the
|
|
||||||
generate-async loop and SSE/JSON response construction.
|
|
||||||
|
|
||||||
Lifecycle::
|
|
||||||
|
|
||||||
handle()
|
|
||||||
├─ build_prompt() # protocol-specific prompt assembly
|
|
||||||
├─ create_response_id() # unique response identifier
|
|
||||||
├─ [stream]
|
|
||||||
│ ├─ format_stream_start()
|
|
||||||
│ ├─ format_stream_token() × N
|
|
||||||
│ │ └─ on_token() hook for stop-sequence interception
|
|
||||||
│ └─ format_stream_end()
|
|
||||||
└─ [non-stream]
|
|
||||||
├─ (accumulate tokens)
|
|
||||||
└─ format_non_stream_response()
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
request_model: type[BaseModel]
|
@abstractmethod
|
||||||
|
def prepare(
|
||||||
|
self, request: BaseModel, engine: InferenceEngine
|
||||||
|
) -> Tuple[str, GenContext, List[str]]:
|
||||||
|
"""Return (prompt, ctx, stop_sequences) for a generation request."""
|
||||||
|
|
||||||
def __init__(self, request: BaseModel, engine: InferenceEngine):
|
@abstractmethod
|
||||||
|
def format_stream_start(self, ctx: GenContext) -> List[str]:
|
||||||
|
"""SSE events that open the stream."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def format_chunk(self, token: str) -> str:
|
||||||
|
"""SSE event for a single generated token."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||||
|
"""SSE events that close the stream."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def format_response(
|
||||||
|
self, ctx: GenContext, content: str, stop: StopInfo
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""JSON response body for non-streaming mode."""
|
||||||
|
|
||||||
|
|
||||||
|
class ProtocolHandler:
|
||||||
|
"""Orchestrates the generation loop, delegates formatting to a builder.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
handler = ProtocolHandler(request, engine, OpenAIResponseBuilder())
|
||||||
|
response = await handler.handle()
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, request: BaseModel, engine: InferenceEngine, builder: ResponseBuilder
|
||||||
|
):
|
||||||
self.request = request
|
self.request = request
|
||||||
self.engine = engine
|
self.engine = engine
|
||||||
|
self.builder = builder
|
||||||
@abstractmethod
|
|
||||||
def build_prompt(self) -> str:
|
|
||||||
"""Build the full prompt string from the request messages."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def create_response_id(self) -> str:
|
|
||||||
"""Generate a unique response ID following the protocol convention."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
|
||||||
"""Yield SSE events that open the stream (role marker, metadata)."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
|
||||||
"""Yield an SSE event for a single generated token."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
|
||||||
"""Yield SSE events that close the stream (finish reason, usage stats)."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def format_non_stream_response(
|
|
||||||
self, ctx: StreamContext, content: str
|
|
||||||
) -> Dict[str, Any]:
|
|
||||||
"""Build the JSON response body for non-streaming mode."""
|
|
||||||
|
|
||||||
def get_stop_sequences(self) -> List[str]:
|
|
||||||
return []
|
|
||||||
|
|
||||||
def create_stop_checker(self) -> StopChecker:
|
|
||||||
return StopChecker(self.get_stop_sequences())
|
|
||||||
|
|
||||||
def on_token(
|
|
||||||
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
|
||||||
) -> Optional[str]:
|
|
||||||
"""Hook after each token is appended to accumulated.
|
|
||||||
|
|
||||||
Return a matched stop-sequence string to break the loop,
|
|
||||||
or None to continue.
|
|
||||||
|
|
||||||
"""
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def handle(self) -> Union[StreamingResponse, Dict[str, Any]]:
|
async def handle(self) -> Union[StreamingResponse, Dict[str, Any]]:
|
||||||
ctx = StreamContext(
|
prompt, ctx, stop_sequences = self.builder.prepare(self.request, self.engine)
|
||||||
resp_id=self.create_response_id(),
|
ctx.prompt_tokens = len(self.engine.tokenizer.encode(prompt))
|
||||||
created=int(time.time()),
|
|
||||||
model=self.request.model,
|
|
||||||
prompt_tokens=self._count_prompt_tokens(),
|
|
||||||
)
|
|
||||||
|
|
||||||
agen = self.engine.generate_async(
|
agen = self.engine.generate_async(
|
||||||
prompt=self.build_prompt(),
|
prompt=prompt,
|
||||||
max_tokens=self.request.max_tokens,
|
max_tokens=self.request.max_tokens,
|
||||||
temperature=self.request.temperature,
|
temperature=self.request.temperature,
|
||||||
top_p=self.request.top_p,
|
top_p=self.request.top_p,
|
||||||
@@ -152,33 +121,37 @@ class ProtocolHandler(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.request.stream:
|
if self.request.stream:
|
||||||
return self._handle_stream(agen, ctx)
|
return self._handle_stream(agen, ctx, stop_sequences)
|
||||||
else:
|
else:
|
||||||
return await self._handle_non_stream(agen, ctx)
|
return await self._handle_non_stream(agen, ctx, stop_sequences)
|
||||||
|
|
||||||
def _count_prompt_tokens(self) -> int:
|
def _handle_stream(
|
||||||
return len(self.engine.tokenizer.encode(self.build_prompt()))
|
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||||
|
) -> StreamingResponse:
|
||||||
def _handle_stream(self, agen, ctx: StreamContext) -> StreamingResponse:
|
checker = StopChecker(stop_sequences)
|
||||||
stop_checker = self.create_stop_checker()
|
|
||||||
|
|
||||||
async def event_stream():
|
async def event_stream():
|
||||||
for event in self.format_stream_start(ctx):
|
for event in self.builder.format_stream_start(ctx):
|
||||||
yield event
|
yield event
|
||||||
|
|
||||||
|
body = ""
|
||||||
|
yielded = ""
|
||||||
|
matched = None
|
||||||
async for token in agen:
|
async for token in agen:
|
||||||
ctx.completion_tokens += 1
|
ctx.completion_tokens += 1
|
||||||
ctx.accumulated += token
|
body += token
|
||||||
|
|
||||||
matched = self.on_token(ctx, token, stop_checker)
|
matched = checker.check(body)
|
||||||
if matched:
|
if matched:
|
||||||
break
|
break
|
||||||
|
|
||||||
yield self.format_stream_token(ctx, token)
|
yield self.builder.format_chunk(token)
|
||||||
|
yielded += token
|
||||||
|
|
||||||
for event in self.format_stream_end(ctx):
|
stop = StopInfo(matched=matched, body=body, yielded=yielded)
|
||||||
|
for event in self.builder.format_stream_end(ctx, stop):
|
||||||
yield event
|
yield event
|
||||||
yield _sse_done()
|
yield sse_done()
|
||||||
|
|
||||||
return StreamingResponse(
|
return StreamingResponse(
|
||||||
event_stream(),
|
event_stream(),
|
||||||
@@ -186,260 +159,23 @@ class ProtocolHandler(ABC):
|
|||||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _handle_non_stream(self, agen, ctx: StreamContext) -> Dict[str, Any]:
|
async def _handle_non_stream(
|
||||||
stop_checker = self.create_stop_checker()
|
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
checker = StopChecker(stop_sequences)
|
||||||
chunks: List[str] = []
|
chunks: List[str] = []
|
||||||
|
body = ""
|
||||||
|
matched = None
|
||||||
|
|
||||||
async for token in agen:
|
async for token in agen:
|
||||||
ctx.completion_tokens += 1
|
ctx.completion_tokens += 1
|
||||||
ctx.accumulated += token
|
|
||||||
chunks.append(token)
|
chunks.append(token)
|
||||||
|
body += token
|
||||||
|
|
||||||
matched = self.on_token(ctx, token, stop_checker)
|
matched = checker.check(body)
|
||||||
if matched:
|
if matched:
|
||||||
break
|
break
|
||||||
|
|
||||||
content = "".join(chunks)
|
content = "".join(chunks)
|
||||||
return self.format_non_stream_response(ctx, content)
|
stop = StopInfo(matched=matched, body=body)
|
||||||
|
return self.builder.format_response(ctx, content, stop)
|
||||||
|
|
||||||
def _extract_text_content(content: Union[str, List[Dict[str, Any]]]) -> str:
|
|
||||||
"""Extract plain text from an Anthropic content block (string or list)."""
|
|
||||||
if isinstance(content, str):
|
|
||||||
return content
|
|
||||||
if isinstance(content, list):
|
|
||||||
for block in content:
|
|
||||||
if isinstance(block, dict) and block.get("type") == "text":
|
|
||||||
return block.get("text", "")
|
|
||||||
return ""
|
|
||||||
|
|
||||||
|
|
||||||
class OpenAIHandler(ProtocolHandler):
|
|
||||||
"""OpenAI-compatible /v1/chat/completions handler."""
|
|
||||||
|
|
||||||
def build_prompt(self) -> str:
|
|
||||||
messages = [
|
|
||||||
{"role": m.role, "content": m.content} for m in self.request.messages
|
|
||||||
]
|
|
||||||
return self.engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
|
||||||
|
|
||||||
def create_response_id(self) -> str:
|
|
||||||
return f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
|
||||||
|
|
||||||
def get_stop_sequences(self) -> List[str]:
|
|
||||||
stop = self.request.stop
|
|
||||||
if stop is None:
|
|
||||||
return []
|
|
||||||
return [stop] if isinstance(stop, str) else stop
|
|
||||||
|
|
||||||
def on_token(
|
|
||||||
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
|
||||||
) -> Optional[str]:
|
|
||||||
return stop_checker.check(ctx.accumulated)
|
|
||||||
|
|
||||||
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
|
||||||
return [
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"object": "chat.completion.chunk",
|
|
||||||
"created": ctx.created,
|
|
||||||
"model": ctx.model,
|
|
||||||
"choices": [
|
|
||||||
{
|
|
||||||
"index": 0,
|
|
||||||
"delta": {"role": "assistant"},
|
|
||||||
"finish_reason": None,
|
|
||||||
}
|
|
||||||
],
|
|
||||||
}
|
|
||||||
)
|
|
||||||
]
|
|
||||||
|
|
||||||
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
|
||||||
return _sse_event(
|
|
||||||
{
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"object": "chat.completion.chunk",
|
|
||||||
"created": ctx.created,
|
|
||||||
"model": ctx.model,
|
|
||||||
"choices": [
|
|
||||||
{"index": 0, "delta": {"content": token}, "finish_reason": None}
|
|
||||||
],
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
|
||||||
return [
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"object": "chat.completion.chunk",
|
|
||||||
"created": ctx.created,
|
|
||||||
"model": ctx.model,
|
|
||||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
|
||||||
}
|
|
||||||
),
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"prompt_tokens": ctx.prompt_tokens,
|
|
||||||
"completion_tokens": ctx.completion_tokens,
|
|
||||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
|
||||||
}
|
|
||||||
),
|
|
||||||
]
|
|
||||||
|
|
||||||
def format_non_stream_response(
|
|
||||||
self, ctx: StreamContext, content: str
|
|
||||||
) -> Dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"object": "chat.completion",
|
|
||||||
"created": ctx.created,
|
|
||||||
"model": ctx.model,
|
|
||||||
"choices": [
|
|
||||||
{
|
|
||||||
"index": 0,
|
|
||||||
"message": {"role": "assistant", "content": content},
|
|
||||||
"finish_reason": "stop",
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"usage": {
|
|
||||||
"prompt_tokens": ctx.prompt_tokens,
|
|
||||||
"completion_tokens": ctx.completion_tokens,
|
|
||||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class AnthropicHandler(ProtocolHandler):
|
|
||||||
"""Anthropic-compatible /v1/messages handler."""
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self._yielded = ""
|
|
||||||
|
|
||||||
def build_prompt(self) -> str:
|
|
||||||
messages: List[Dict[str, str]] = []
|
|
||||||
system = getattr(self.request, "system", None)
|
|
||||||
if system:
|
|
||||||
messages.append({"role": "system", "content": system})
|
|
||||||
for m in self.request.messages:
|
|
||||||
content = _extract_text_content(m.content)
|
|
||||||
if content:
|
|
||||||
messages.append({"role": m.role, "content": content})
|
|
||||||
return self.engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
|
||||||
|
|
||||||
def create_response_id(self) -> str:
|
|
||||||
return f"msg_{uuid.uuid4().hex[:24]}"
|
|
||||||
|
|
||||||
def get_stop_sequences(self) -> List[str]:
|
|
||||||
return getattr(self.request, "stop_sequences", None) or []
|
|
||||||
|
|
||||||
def on_token(
|
|
||||||
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
|
||||||
) -> Optional[str]:
|
|
||||||
matched = stop_checker.check(ctx.accumulated)
|
|
||||||
if not matched:
|
|
||||||
return None
|
|
||||||
|
|
||||||
ctx.stop_matched = matched
|
|
||||||
trimmed = ctx.accumulated[: ctx.accumulated.rfind(matched)]
|
|
||||||
unyielded = trimmed[len(self._yielded) :]
|
|
||||||
if unyielded:
|
|
||||||
ctx.last_yield_trimmed = unyielded
|
|
||||||
return matched
|
|
||||||
|
|
||||||
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
|
||||||
return [
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"type": "message_start",
|
|
||||||
"message": {
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"type": "message",
|
|
||||||
"role": "assistant",
|
|
||||||
"model": ctx.model,
|
|
||||||
"content": [],
|
|
||||||
"usage": {"input_tokens": ctx.prompt_tokens},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
event="message_start",
|
|
||||||
),
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"type": "content_block_start",
|
|
||||||
"index": 0,
|
|
||||||
"content_block": {"type": "text", "text": ""},
|
|
||||||
},
|
|
||||||
event="content_block_start",
|
|
||||||
),
|
|
||||||
]
|
|
||||||
|
|
||||||
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
|
||||||
self._yielded += token
|
|
||||||
return _sse_event(
|
|
||||||
{
|
|
||||||
"type": "content_block_delta",
|
|
||||||
"index": 0,
|
|
||||||
"delta": {"type": "text_delta", "text": token},
|
|
||||||
},
|
|
||||||
event="content_block_delta",
|
|
||||||
)
|
|
||||||
|
|
||||||
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
|
||||||
matched = ctx.stop_matched
|
|
||||||
events: List[str] = []
|
|
||||||
last_yielded = ctx.last_yield_trimmed
|
|
||||||
if last_yielded:
|
|
||||||
events.append(
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"type": "content_block_delta",
|
|
||||||
"index": 0,
|
|
||||||
"delta": {"type": "text_delta", "text": last_yielded},
|
|
||||||
},
|
|
||||||
event="content_block_delta",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
events.append(
|
|
||||||
_sse_event(
|
|
||||||
{"type": "content_block_stop", "index": 0},
|
|
||||||
event="content_block_stop",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
events.append(
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"type": "message_delta",
|
|
||||||
"delta": {
|
|
||||||
"stop_reason": "stop_sequence" if matched else "end_turn",
|
|
||||||
"stop_sequence": matched,
|
|
||||||
},
|
|
||||||
"usage": {"output_tokens": ctx.completion_tokens},
|
|
||||||
},
|
|
||||||
event="message_delta",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
events.append(_sse_event({"type": "message_stop"}, event="message_stop"))
|
|
||||||
return events
|
|
||||||
|
|
||||||
def format_non_stream_response(
|
|
||||||
self, ctx: StreamContext, content: str
|
|
||||||
) -> Dict[str, Any]:
|
|
||||||
matched = ctx.stop_matched
|
|
||||||
if matched:
|
|
||||||
content = content[: content.rfind(matched)]
|
|
||||||
return {
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"type": "message",
|
|
||||||
"role": "assistant",
|
|
||||||
"model": ctx.model,
|
|
||||||
"content": [{"type": "text", "text": content}],
|
|
||||||
"stop_reason": "stop_sequence" if matched else "end_turn",
|
|
||||||
"stop_sequence": matched,
|
|
||||||
"usage": {
|
|
||||||
"input_tokens": ctx.prompt_tokens,
|
|
||||||
"output_tokens": ctx.completion_tokens,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -15,7 +15,9 @@ import uvicorn
|
|||||||
from fastapi import FastAPI, HTTPException
|
from fastapi import FastAPI, HTTPException
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from astrai.inference.api.protocol import AnthropicHandler, OpenAIHandler
|
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||||
|
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||||
|
from astrai.inference.api.protocol import ProtocolHandler
|
||||||
from astrai.inference.engine import InferenceEngine
|
from astrai.inference.engine import InferenceEngine
|
||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
@@ -133,14 +135,14 @@ async def get_stats():
|
|||||||
@app.post("/v1/chat/completions")
|
@app.post("/v1/chat/completions")
|
||||||
async def chat_completion(request: ChatCompletionRequest):
|
async def chat_completion(request: ChatCompletionRequest):
|
||||||
engine = _get_engine()
|
engine = _get_engine()
|
||||||
handler = OpenAIHandler(request, engine)
|
handler = ProtocolHandler(request, engine, OpenAIResponseBuilder())
|
||||||
return await handler.handle()
|
return await handler.handle()
|
||||||
|
|
||||||
|
|
||||||
@app.post("/v1/messages")
|
@app.post("/v1/messages")
|
||||||
async def create_message(request: MessagesRequest):
|
async def create_message(request: MessagesRequest):
|
||||||
engine = _get_engine()
|
engine = _get_engine()
|
||||||
handler = AnthropicHandler(request, engine)
|
handler = ProtocolHandler(request, engine, AnthropicResponseBuilder())
|
||||||
return await handler.handle()
|
return await handler.handle()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ class Allocator:
|
|||||||
return idx
|
return idx
|
||||||
return -1
|
return -1
|
||||||
|
|
||||||
def free(self, idx: int, keep_cached: bool = False) -> None:
|
def free(self, idx: int, keep_cached: bool = False):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._refs[idx] -= 1
|
self._refs[idx] -= 1
|
||||||
if self._refs[idx] == 0:
|
if self._refs[idx] == 0:
|
||||||
@@ -51,7 +51,7 @@ class Allocator:
|
|||||||
else:
|
else:
|
||||||
self._free_mask |= 1 << idx
|
self._free_mask |= 1 << idx
|
||||||
|
|
||||||
def inc_ref(self, idx: int) -> None:
|
def inc_ref(self, idx: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._refs[idx] += 1
|
self._refs[idx] += 1
|
||||||
self._lru.pop(idx, None)
|
self._lru.pop(idx, None)
|
||||||
@@ -60,7 +60,7 @@ class Allocator:
|
|||||||
with self._lock:
|
with self._lock:
|
||||||
return self._refs[idx]
|
return self._refs[idx]
|
||||||
|
|
||||||
def touch(self, idx: int) -> None:
|
def touch(self, idx: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._lru.move_to_end(idx)
|
self._lru.move_to_end(idx)
|
||||||
|
|
||||||
@@ -74,7 +74,7 @@ class PrefixCache:
|
|||||||
self._hash_to_page: Dict[int, int] = {}
|
self._hash_to_page: Dict[int, int] = {}
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
def evict(self, idx: int) -> None:
|
def evict(self, idx: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
h = self._page_to_hash.pop(idx, None)
|
h = self._page_to_hash.pop(idx, None)
|
||||||
if h is not None:
|
if h is not None:
|
||||||
@@ -96,9 +96,7 @@ class PrefixCache:
|
|||||||
hits.append(p)
|
hits.append(p)
|
||||||
return hits
|
return hits
|
||||||
|
|
||||||
def record(
|
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
||||||
self, page_idx: int, token_ids: List[int], logical_page_idx: int
|
|
||||||
) -> None:
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
h = page_hash(token_ids, logical_page_idx, self._page_size)
|
h = page_hash(token_ids, logical_page_idx, self._page_size)
|
||||||
old_h = self._page_to_hash.pop(page_idx, None)
|
old_h = self._page_to_hash.pop(page_idx, None)
|
||||||
@@ -127,13 +125,13 @@ class PagePool:
|
|||||||
def alloc(self) -> int:
|
def alloc(self) -> int:
|
||||||
return self._alloc.alloc()
|
return self._alloc.alloc()
|
||||||
|
|
||||||
def free(self, idx: int) -> None:
|
def free(self, idx: int):
|
||||||
keep = self._prefix.has_page(idx)
|
keep = self._prefix.has_page(idx)
|
||||||
self._alloc.free(idx, keep_cached=keep)
|
self._alloc.free(idx, keep_cached=keep)
|
||||||
if not keep:
|
if not keep:
|
||||||
self._prefix.evict(idx)
|
self._prefix.evict(idx)
|
||||||
|
|
||||||
def inc_ref(self, idx: int) -> None:
|
def inc_ref(self, idx: int):
|
||||||
self._alloc.inc_ref(idx)
|
self._alloc.inc_ref(idx)
|
||||||
|
|
||||||
def lookup(self, token_ids: List[int]) -> List[int]:
|
def lookup(self, token_ids: List[int]) -> List[int]:
|
||||||
@@ -142,9 +140,7 @@ class PagePool:
|
|||||||
self._alloc.touch(p)
|
self._alloc.touch(p)
|
||||||
return hits
|
return hits
|
||||||
|
|
||||||
def record(
|
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
||||||
self, page_idx: int, token_ids: List[int], logical_page_idx: int
|
|
||||||
) -> None:
|
|
||||||
self._prefix.record(page_idx, token_ids, logical_page_idx)
|
self._prefix.record(page_idx, token_ids, logical_page_idx)
|
||||||
|
|
||||||
|
|
||||||
@@ -157,7 +153,7 @@ class TaskTable:
|
|||||||
self._cached: Dict[str, int] = {}
|
self._cached: Dict[str, int] = {}
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
def set(self, task_id: str, page_table: List[int], cached: int) -> None:
|
def set(self, task_id: str, page_table: List[int], cached: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._pages[task_id] = page_table
|
self._pages[task_id] = page_table
|
||||||
self._cached[task_id] = cached
|
self._cached[task_id] = cached
|
||||||
@@ -220,7 +216,7 @@ class Storage:
|
|||||||
start_pos: int,
|
start_pos: int,
|
||||||
k: Tensor,
|
k: Tensor,
|
||||||
v: Tensor,
|
v: Tensor,
|
||||||
) -> None:
|
):
|
||||||
seq_len = k.size(1)
|
seq_len = k.size(1)
|
||||||
if seq_len == 0:
|
if seq_len == 0:
|
||||||
return
|
return
|
||||||
@@ -286,7 +282,7 @@ class KvcacheView:
|
|||||||
self._page_table = page_table
|
self._page_table = page_table
|
||||||
self._total_len = total_len
|
self._total_len = total_len
|
||||||
|
|
||||||
def write(self, layer_id: int, k: Tensor, v: Tensor) -> None:
|
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
||||||
start_pos = self._total_len - k.size(1)
|
start_pos = self._total_len - k.size(1)
|
||||||
self._storage.write(layer_id, self._page_table, start_pos, k, v)
|
self._storage.write(layer_id, self._page_table, start_pos, k, v)
|
||||||
|
|
||||||
@@ -339,7 +335,7 @@ class KVCache:
|
|||||||
self._table.set(task_id, hits + new_pages, cached)
|
self._table.set(task_id, hits + new_pages, cached)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def task_free(self, task_id: str) -> None:
|
def task_free(self, task_id: str):
|
||||||
page_table, _ = self._table.pop(task_id)
|
page_table, _ = self._table.pop(task_id)
|
||||||
for idx in page_table:
|
for idx in page_table:
|
||||||
self._pool.free(idx)
|
self._pool.free(idx)
|
||||||
@@ -359,7 +355,7 @@ class KVCache:
|
|||||||
|
|
||||||
def task_record_hashes(
|
def task_record_hashes(
|
||||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||||
) -> None:
|
):
|
||||||
page_table = self._table.get(task_id)
|
page_table = self._table.get(task_id)
|
||||||
full_pages = len(prompt_ids) // self.page_size
|
full_pages = len(prompt_ids) // self.page_size
|
||||||
for i in range(start_logical_page, full_pages):
|
for i in range(start_logical_page, full_pages):
|
||||||
|
|||||||
@@ -29,9 +29,7 @@ class Executor:
|
|||||||
self.device = device or next(model.parameters()).device
|
self.device = device or next(model.parameters()).device
|
||||||
self.dtype = dtype or next(model.parameters()).dtype
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
def execute_prefill(
|
def execute_prefill(self, tasks: List[Task], prompt_len: int, start_pos: int = 0):
|
||||||
self, tasks: List[Task], prompt_len: int, start_pos: int = 0
|
|
||||||
) -> None:
|
|
||||||
if start_pos >= prompt_len:
|
if start_pos >= prompt_len:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|||||||
@@ -75,14 +75,14 @@ class InferenceScheduler:
|
|||||||
def add_task(self, prompt: str, **kwargs) -> str:
|
def add_task(self, prompt: str, **kwargs) -> str:
|
||||||
return self._task_mgr.add_task(prompt, **kwargs)
|
return self._task_mgr.add_task(prompt, **kwargs)
|
||||||
|
|
||||||
def remove_task(self, task_id: str) -> None:
|
def remove_task(self, task_id: str):
|
||||||
for task in self._task_mgr.remove_task(task_id):
|
for task in self._task_mgr.remove_task(task_id):
|
||||||
self._page_cache.task_free(task.task_id)
|
self._page_cache.task_free(task.task_id)
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return self._task_mgr.get_stats()
|
return self._task_mgr.get_stats()
|
||||||
|
|
||||||
def _run_generation_loop(self) -> None:
|
def _run_generation_loop(self):
|
||||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||||
try:
|
try:
|
||||||
while self._running:
|
while self._running:
|
||||||
@@ -108,7 +108,10 @@ class InferenceScheduler:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
to_prefill = [
|
to_prefill = [
|
||||||
t for t in self._task_mgr.get_active_tasks() if t.output_tokens == 0
|
t
|
||||||
|
for t in self._task_mgr.get_active_tasks()
|
||||||
|
if t.output_tokens == 0
|
||||||
|
and self._page_cache.task_cached(t.task_id) < len(t.prompt_ids)
|
||||||
]
|
]
|
||||||
if to_prefill:
|
if to_prefill:
|
||||||
for t in to_prefill:
|
for t in to_prefill:
|
||||||
@@ -183,14 +186,14 @@ class InferenceScheduler:
|
|||||||
self._task_mgr.clear_queues()
|
self._task_mgr.clear_queues()
|
||||||
raise
|
raise
|
||||||
|
|
||||||
def start(self) -> None:
|
def start(self):
|
||||||
if not self._running:
|
if not self._running:
|
||||||
self._running = True
|
self._running = True
|
||||||
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||||
t.start()
|
t.start()
|
||||||
self._loop_thread = t
|
self._loop_thread = t
|
||||||
|
|
||||||
def stop(self) -> None:
|
def stop(self):
|
||||||
self._running = False
|
self._running = False
|
||||||
self._task_mgr.wake()
|
self._task_mgr.wake()
|
||||||
if hasattr(self, "_loop_thread"):
|
if hasattr(self, "_loop_thread"):
|
||||||
|
|||||||
@@ -172,12 +172,12 @@ class TaskManager:
|
|||||||
to_add.append(self.waiting_queue.popleft())
|
to_add.append(self.waiting_queue.popleft())
|
||||||
return to_add
|
return to_add
|
||||||
|
|
||||||
def activate(self, task: Task) -> None:
|
def activate(self, task: Task):
|
||||||
task.status = TaskStatus.RUNNING
|
task.status = TaskStatus.RUNNING
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.active_tasks.append(task)
|
self.active_tasks.append(task)
|
||||||
|
|
||||||
def return_to_waiting(self, tasks: List[Task]) -> None:
|
def return_to_waiting(self, tasks: List[Task]):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
for task in reversed(tasks):
|
for task in reversed(tasks):
|
||||||
self.waiting_queue.appendleft(task)
|
self.waiting_queue.appendleft(task)
|
||||||
@@ -185,7 +185,7 @@ class TaskManager:
|
|||||||
def has_work(self) -> bool:
|
def has_work(self) -> bool:
|
||||||
return bool(self.active_tasks or self.waiting_queue)
|
return bool(self.active_tasks or self.waiting_queue)
|
||||||
|
|
||||||
def wait_for_tasks(self, timeout: float = 1.0) -> None:
|
def wait_for_tasks(self, timeout: float = 1.0):
|
||||||
self._task_event.clear()
|
self._task_event.clear()
|
||||||
self._task_event.wait(timeout=timeout)
|
self._task_event.wait(timeout=timeout)
|
||||||
|
|
||||||
@@ -197,10 +197,10 @@ class TaskManager:
|
|||||||
with self._lock:
|
with self._lock:
|
||||||
return list(self.waiting_queue)
|
return list(self.waiting_queue)
|
||||||
|
|
||||||
def clear_queues(self) -> None:
|
def clear_queues(self):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.waiting_queue.clear()
|
self.waiting_queue.clear()
|
||||||
self.active_tasks.clear()
|
self.active_tasks.clear()
|
||||||
|
|
||||||
def wake(self) -> None:
|
def wake(self):
|
||||||
self._task_event.set()
|
self._task_event.set()
|
||||||
|
|||||||
@@ -13,17 +13,6 @@ from astrai.inference.core.task import STOP
|
|||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
def _validate_sampling_params(
|
|
||||||
top_k: int, top_p: float, temperature: float, max_tokens: Optional[int] = None
|
|
||||||
):
|
|
||||||
if not (isinstance(top_k, int) and top_k >= 0):
|
|
||||||
raise ValueError("top_k must be a non-negative integer")
|
|
||||||
if not (0.0 <= top_p <= 1.0):
|
|
||||||
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
|
||||||
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
|
||||||
raise ValueError("temperature must be a non-negative number")
|
|
||||||
|
|
||||||
|
|
||||||
class GenerateResult:
|
class GenerateResult:
|
||||||
"""Thread-safe token accumulator for streaming and non-streaming modes."""
|
"""Thread-safe token accumulator for streaming and non-streaming modes."""
|
||||||
|
|
||||||
@@ -59,7 +48,7 @@ class GenerateResult:
|
|||||||
def wait(self, timeout: Optional[float] = None) -> bool:
|
def wait(self, timeout: Optional[float] = None) -> bool:
|
||||||
return self._event.wait(timeout=timeout)
|
return self._event.wait(timeout=timeout)
|
||||||
|
|
||||||
def wait_completion(self, timeout: float = 300.0) -> None:
|
def wait_completion(self, timeout: float = 300.0):
|
||||||
with self._cond:
|
with self._cond:
|
||||||
if not self._cond.wait_for(
|
if not self._cond.wait_for(
|
||||||
lambda: self._completed >= self._total, timeout=timeout
|
lambda: self._completed >= self._total, timeout=timeout
|
||||||
@@ -86,7 +75,12 @@ class GenerationRequest:
|
|||||||
max_tokens: Optional[int] = None,
|
max_tokens: Optional[int] = None,
|
||||||
stream: bool = False,
|
stream: bool = False,
|
||||||
):
|
):
|
||||||
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
if not (isinstance(top_k, int) and top_k >= 0):
|
||||||
|
raise ValueError("top_k must be a non-negative integer")
|
||||||
|
if not (0.0 <= top_p <= 1.0):
|
||||||
|
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||||
|
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
||||||
|
raise ValueError("temperature must be a non-negative number")
|
||||||
|
|
||||||
self.messages = messages
|
self.messages = messages
|
||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
@@ -137,7 +131,6 @@ class InferenceEngine:
|
|||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
) -> Union[Generator, str, List[str]]:
|
) -> Union[Generator, str, List[str]]:
|
||||||
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
|
||||||
is_batch = isinstance(prompt, list)
|
is_batch = isinstance(prompt, list)
|
||||||
prompts = prompt if is_batch else [prompt]
|
prompts = prompt if is_batch else [prompt]
|
||||||
|
|
||||||
@@ -158,7 +151,6 @@ class InferenceEngine:
|
|||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
|
||||||
sync_gen = self._generate_streaming(
|
sync_gen = self._generate_streaming(
|
||||||
[prompt], False, max_tokens, temperature, top_p, top_k
|
[prompt], False, max_tokens, temperature, top_p, top_k
|
||||||
)
|
)
|
||||||
@@ -289,7 +281,7 @@ class InferenceEngine:
|
|||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return self.scheduler.get_stats()
|
return self.scheduler.get_stats()
|
||||||
|
|
||||||
def shutdown(self) -> None:
|
def shutdown(self):
|
||||||
self.scheduler.stop()
|
self.scheduler.stop()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|||||||
@@ -2,6 +2,13 @@ from astrai.model.automodel import AutoModel
|
|||||||
from astrai.model.components.attention import GQA
|
from astrai.model.components.attention import GQA
|
||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.model.components.linear import Linear
|
from astrai.model.components.linear import Linear
|
||||||
|
from astrai.model.components.lora import (
|
||||||
|
LoRAConfig,
|
||||||
|
inject_lora,
|
||||||
|
load_lora,
|
||||||
|
merge_lora,
|
||||||
|
save_lora,
|
||||||
|
)
|
||||||
from astrai.model.components.mlp import MLP
|
from astrai.model.components.mlp import MLP
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
from astrai.model.encoder import EmbeddingEncoder
|
from astrai.model.encoder import EmbeddingEncoder
|
||||||
@@ -18,4 +25,10 @@ __all__ = [
|
|||||||
"AutoRegressiveLM",
|
"AutoRegressiveLM",
|
||||||
"EmbeddingEncoder",
|
"EmbeddingEncoder",
|
||||||
"AutoModel",
|
"AutoModel",
|
||||||
|
# LoRA
|
||||||
|
"LoRAConfig",
|
||||||
|
"inject_lora",
|
||||||
|
"merge_lora",
|
||||||
|
"save_lora",
|
||||||
|
"load_lora",
|
||||||
]
|
]
|
||||||
|
|||||||
+23
-29
@@ -2,21 +2,24 @@
|
|||||||
AutoModel base class for model loading and saving.
|
AutoModel base class for model loading and saving.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Self, Union
|
from typing import Self, Union
|
||||||
|
|
||||||
import safetensors.torch as st
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from astrai.config.model_config import BaseModelConfig, ConfigFactory
|
from astrai.config.model_config import BaseModelConfig, ConfigFactory
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.serialization import load_model_config, load_model_weights, save_model
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def _disable_random_init(enable: bool = True):
|
def _disable_random_init(enable: bool = True):
|
||||||
init_functions = [
|
if not enable:
|
||||||
|
yield
|
||||||
|
return
|
||||||
|
|
||||||
|
names = (
|
||||||
"xavier_normal_",
|
"xavier_normal_",
|
||||||
"xavier_uniform_",
|
"xavier_uniform_",
|
||||||
"kaiming_normal_",
|
"kaiming_normal_",
|
||||||
@@ -26,18 +29,15 @@ def _disable_random_init(enable: bool = True):
|
|||||||
"constant_",
|
"constant_",
|
||||||
"normal_",
|
"normal_",
|
||||||
"uniform_",
|
"uniform_",
|
||||||
]
|
)
|
||||||
original_funcs = {}
|
orig = {n: getattr(nn.init, n) for n in names if hasattr(nn.init, n)}
|
||||||
for name in init_functions:
|
for n in orig:
|
||||||
if enable and hasattr(nn.init, name):
|
setattr(nn.init, n, lambda *a, **kw: None)
|
||||||
original_funcs[name] = getattr(nn.init, name)
|
|
||||||
setattr(nn.init, name, lambda *args, **kwargs: None)
|
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
finally:
|
finally:
|
||||||
if enable:
|
for n, fn in orig.items():
|
||||||
for name, orig_func in original_funcs.items():
|
setattr(nn.init, n, fn)
|
||||||
setattr(nn.init, name, orig_func)
|
|
||||||
|
|
||||||
|
|
||||||
class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
||||||
@@ -60,25 +60,22 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
|||||||
|
|
||||||
model_path = Path(path)
|
model_path = Path(path)
|
||||||
|
|
||||||
# Load config
|
|
||||||
config_path = model_path / "config.json"
|
config_path = model_path / "config.json"
|
||||||
if config_path.exists():
|
if not config_path.exists():
|
||||||
with open(config_path, "r") as f:
|
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||||
raw = json.load(f)
|
|
||||||
|
raw = load_model_config(str(model_path))
|
||||||
config = ConfigFactory.load(raw)
|
config = ConfigFactory.load(raw)
|
||||||
model_type = config.model_type or "autoregressive_lm"
|
model_type = config.model_type or "autoregressive_lm"
|
||||||
else:
|
|
||||||
raise FileNotFoundError(f"Config file not found: {config_path}")
|
|
||||||
|
|
||||||
actual_cls = AutoModel.get_component_class(model_type)
|
actual_cls = AutoModel.get_component_class(model_type)
|
||||||
|
|
||||||
with _disable_random_init(enable=disable_random_init):
|
with _disable_random_init(enable=disable_random_init):
|
||||||
model = actual_cls(config)
|
model = actual_cls(config)
|
||||||
|
|
||||||
# Load weights
|
|
||||||
weights_path = model_path / "model.safetensors"
|
weights_path = model_path / "model.safetensors"
|
||||||
if weights_path.exists():
|
if weights_path.exists():
|
||||||
state_dict = st.load_file(str(weights_path))
|
state_dict = load_model_weights(str(model_path))
|
||||||
model.load_state_dict(state_dict, strict=strict)
|
model.load_state_dict(state_dict, strict=strict)
|
||||||
|
|
||||||
return model
|
return model
|
||||||
@@ -86,15 +83,12 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
|||||||
def save_pretrained(
|
def save_pretrained(
|
||||||
self,
|
self,
|
||||||
save_directory: Union[str, Path],
|
save_directory: Union[str, Path],
|
||||||
) -> None:
|
):
|
||||||
save_path = Path(save_directory)
|
save_model(
|
||||||
save_path.mkdir(parents=True, exist_ok=True)
|
config=self.config.to_dict(),
|
||||||
|
state_dict=self.state_dict(),
|
||||||
# Save config
|
save_directory=str(save_directory),
|
||||||
self.config.to_file(str(save_path / "config.json"))
|
)
|
||||||
|
|
||||||
# Save weights
|
|
||||||
st.save_file(self.state_dict(), str(save_path / "model.safetensors"))
|
|
||||||
|
|
||||||
def to(self, *args, **kwargs) -> Self:
|
def to(self, *args, **kwargs) -> Self:
|
||||||
"""Move model to device/dtype."""
|
"""Move model to device/dtype."""
|
||||||
|
|||||||
@@ -0,0 +1,194 @@
|
|||||||
|
import logging
|
||||||
|
from dataclasses import asdict, dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Optional, Set
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
from astrai.serialization import (
|
||||||
|
load_json,
|
||||||
|
load_safetensors,
|
||||||
|
save_json,
|
||||||
|
save_safetensors,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
TARGET_MODULES_ATTN = {"q_proj", "k_proj", "v_proj", "o_proj"}
|
||||||
|
TARGET_MODULES_FFN = {"up", "gate", "down"}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LoRAConfig:
|
||||||
|
r: int = 16
|
||||||
|
alpha: int = 32
|
||||||
|
target_modules: tuple = ("q_proj", "v_proj")
|
||||||
|
|
||||||
|
|
||||||
|
class LoRALinear(nn.Module):
|
||||||
|
def __init__(self, base: Linear, r: int = 16, alpha: int = 32):
|
||||||
|
super().__init__()
|
||||||
|
self.register_parameter("weight", base.weight)
|
||||||
|
self.weight.requires_grad_(False)
|
||||||
|
self.bias = base.bias
|
||||||
|
if self.bias is not None:
|
||||||
|
self.bias.requires_grad_(False)
|
||||||
|
|
||||||
|
self.r = r
|
||||||
|
self.scaling = alpha / r
|
||||||
|
self.lora_A = nn.Parameter(torch.randn(r, self.weight.shape[1]) / r)
|
||||||
|
self.lora_B = nn.Parameter(torch.zeros(self.weight.shape[0], r))
|
||||||
|
self._merged = False
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
out = F.linear(x, self.weight, self.bias)
|
||||||
|
if not self._merged:
|
||||||
|
out += (F.linear(x, self.lora_A) @ self.lora_B.T) * self.scaling
|
||||||
|
return out
|
||||||
|
|
||||||
|
def merge(self):
|
||||||
|
if self._merged:
|
||||||
|
return
|
||||||
|
self.weight.data += (self.lora_B @ self.lora_A) * self.scaling
|
||||||
|
self._merged = True
|
||||||
|
del self.lora_A
|
||||||
|
del self.lora_B
|
||||||
|
|
||||||
|
|
||||||
|
def _collect_lora_info(model: nn.Module) -> dict:
|
||||||
|
names = {}
|
||||||
|
for n, m in model.named_modules():
|
||||||
|
if isinstance(m, Linear):
|
||||||
|
_, _, child = n.rpartition(".")
|
||||||
|
names.setdefault(child, []).append(n)
|
||||||
|
return names
|
||||||
|
|
||||||
|
|
||||||
|
def _get_lora_count(model: nn.Module) -> int:
|
||||||
|
return sum(1 for m in model.modules() if isinstance(m, LoRALinear))
|
||||||
|
|
||||||
|
|
||||||
|
def inject_lora(
|
||||||
|
model: nn.Module,
|
||||||
|
r: int = 16,
|
||||||
|
alpha: int = 32,
|
||||||
|
target_modules: Optional[Set[str]] = None,
|
||||||
|
) -> LoRAConfig:
|
||||||
|
if target_modules is None:
|
||||||
|
target_modules = TARGET_MODULES_ATTN
|
||||||
|
|
||||||
|
available = _collect_lora_info(model)
|
||||||
|
injected = 0
|
||||||
|
|
||||||
|
for name, module in list(model.named_modules()):
|
||||||
|
if not isinstance(module, Linear):
|
||||||
|
continue
|
||||||
|
parent_name, _, child_name = name.rpartition(".")
|
||||||
|
if child_name not in target_modules:
|
||||||
|
continue
|
||||||
|
parent = model.get_submodule(parent_name) if parent_name else model
|
||||||
|
setattr(parent, child_name, LoRALinear(module, r=r, alpha=alpha))
|
||||||
|
injected += 1
|
||||||
|
|
||||||
|
if injected == 0:
|
||||||
|
logger.warning(
|
||||||
|
"No LoRA layers injected. Available Linear child names: %s. "
|
||||||
|
"target_modules: %s. Check model type and target_modules.",
|
||||||
|
sorted(available),
|
||||||
|
sorted(target_modules),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.info("LoRA injected: %d layers (r=%d, alpha=%d)", injected, r, alpha)
|
||||||
|
|
||||||
|
return LoRAConfig(r=r, alpha=alpha, target_modules=tuple(target_modules))
|
||||||
|
|
||||||
|
|
||||||
|
def merge_lora(model: nn.Module):
|
||||||
|
n = 0
|
||||||
|
for module in model.modules():
|
||||||
|
if isinstance(module, LoRALinear):
|
||||||
|
module.merge()
|
||||||
|
n += 1
|
||||||
|
if n == 0:
|
||||||
|
logger.warning("No LoRA layers to merge.")
|
||||||
|
else:
|
||||||
|
logger.info("Merged %d LoRA layers", n)
|
||||||
|
|
||||||
|
|
||||||
|
def save_lora(model: nn.Module, save_dir: str, config: LoRAConfig):
|
||||||
|
lora_sd = {
|
||||||
|
k: v
|
||||||
|
for k, v in model.state_dict().items()
|
||||||
|
if k.endswith((".lora_A", ".lora_B"))
|
||||||
|
}
|
||||||
|
if not lora_sd:
|
||||||
|
raise RuntimeError(
|
||||||
|
"No LoRA parameters found in model. "
|
||||||
|
"The model may not have been injected or was already merged."
|
||||||
|
)
|
||||||
|
|
||||||
|
path = Path(save_dir)
|
||||||
|
path.mkdir(parents=True, exist_ok=True)
|
||||||
|
save_safetensors(lora_sd, path / "adapter_model.safetensors")
|
||||||
|
save_json(asdict(config), path / "adapter_config.json")
|
||||||
|
logger.info("LoRA adapter saved to %s (%d keys)", save_dir, len(lora_sd))
|
||||||
|
|
||||||
|
|
||||||
|
def load_lora(model: nn.Module, load_dir: str) -> LoRAConfig:
|
||||||
|
path = Path(load_dir)
|
||||||
|
raw = load_json(path / "adapter_config.json")
|
||||||
|
config = LoRAConfig(
|
||||||
|
r=raw["r"], alpha=raw["alpha"], target_modules=tuple(raw["target_modules"])
|
||||||
|
)
|
||||||
|
|
||||||
|
existing = _get_lora_count(model)
|
||||||
|
if existing > 0:
|
||||||
|
logger.warning(
|
||||||
|
"Model already has %d LoRA layers. Skipping injection, "
|
||||||
|
"loading weights onto existing layers only.",
|
||||||
|
existing,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
inject_lora(
|
||||||
|
model,
|
||||||
|
r=config.r,
|
||||||
|
alpha=config.alpha,
|
||||||
|
target_modules=set(config.target_modules),
|
||||||
|
)
|
||||||
|
|
||||||
|
weights = load_safetensors(path / "adapter_model.safetensors")
|
||||||
|
try:
|
||||||
|
missing, unexpected = model.load_state_dict(weights, strict=False)
|
||||||
|
except RuntimeError as e:
|
||||||
|
msg = str(e)
|
||||||
|
if "size mismatch" in msg:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"LoRA weight shapes do not match the model. "
|
||||||
|
f"The adapter config (r={config.r}) may not match the injected layers. "
|
||||||
|
f"Original error: {msg}"
|
||||||
|
) from e
|
||||||
|
raise
|
||||||
|
|
||||||
|
injected = _get_lora_count(model)
|
||||||
|
if injected == 0:
|
||||||
|
raise RuntimeError(
|
||||||
|
"No LoRA layers found after loading. "
|
||||||
|
"Inject LoRA before calling load_lora, or check the adapter config."
|
||||||
|
)
|
||||||
|
|
||||||
|
if missing:
|
||||||
|
lora_missing = [k for k in missing if "lora" in k]
|
||||||
|
if lora_missing:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"LoRA weight keys not found in model: {lora_missing}. "
|
||||||
|
f"The adapter config (r={config.r}) may not match the model."
|
||||||
|
)
|
||||||
|
logger.debug("LoRA load: %d missing base-weight keys (expected)", len(missing))
|
||||||
|
if unexpected:
|
||||||
|
logger.warning("LoRA load: %d unexpected keys", len(unexpected))
|
||||||
|
|
||||||
|
logger.info("LoRA adapter loaded from %s", load_dir)
|
||||||
|
return config
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
from typing import Optional
|
from typing import Dict, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -19,6 +19,10 @@ def get_rotary_emb(
|
|||||||
return torch.complex(cos, sin)
|
return torch.complex(cos, sin)
|
||||||
|
|
||||||
|
|
||||||
|
def ntk_base(base: float, dim: int, factor: float) -> float:
|
||||||
|
return base * (factor ** (dim / (dim - 2)))
|
||||||
|
|
||||||
|
|
||||||
def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
||||||
dtype = x.dtype
|
dtype = x.dtype
|
||||||
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||||
@@ -30,11 +34,25 @@ def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
|||||||
|
|
||||||
|
|
||||||
class RotaryEmbedding(nn.Module):
|
class RotaryEmbedding(nn.Module):
|
||||||
def __init__(self, dim: int, max_len: int, base: float = 10000):
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
max_len: int,
|
||||||
|
base: float = 10000,
|
||||||
|
rope_scaling: Optional[Dict] = None,
|
||||||
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dim = dim
|
self.dim = dim
|
||||||
self.max_len = max_len
|
self.max_len = max_len
|
||||||
self.base = base
|
self.base = base
|
||||||
|
self.rope_scaling = rope_scaling
|
||||||
|
|
||||||
|
if rope_scaling is not None:
|
||||||
|
scaling_type = rope_scaling.get("type", "ntk")
|
||||||
|
factor = rope_scaling.get("factor", 1.0)
|
||||||
|
if scaling_type == "ntk":
|
||||||
|
self.base = ntk_base(base, dim, factor)
|
||||||
|
|
||||||
self._set_rotary_buffer(self.max_len)
|
self._set_rotary_buffer(self.max_len)
|
||||||
|
|
||||||
def _set_rotary_buffer(self, max_len: int):
|
def _set_rotary_buffer(self, max_len: int):
|
||||||
|
|||||||
@@ -20,7 +20,9 @@ class EmbeddingEncoder(AutoModel):
|
|||||||
self.config = config
|
self.config = config
|
||||||
rope_dim = config.dim // config.n_heads
|
rope_dim = config.dim // config.n_heads
|
||||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||||
self.rotary_embedding = RotaryEmbedding(rope_dim, config.max_len, rope_base)
|
self.rotary_embedding = RotaryEmbedding(
|
||||||
|
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||||
|
)
|
||||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||||
|
|
||||||
self.layers = nn.ModuleList(
|
self.layers = nn.ModuleList(
|
||||||
@@ -66,9 +68,6 @@ class EmbeddingEncoder(AutoModel):
|
|||||||
|
|
||||||
x = self.embed_tokens(input_ids)
|
x = self.embed_tokens(input_ids)
|
||||||
|
|
||||||
if position_ids is None:
|
|
||||||
position_ids = torch.arange(S, device=x.device).unsqueeze(0).expand(B, -1)
|
|
||||||
|
|
||||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||||
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=False)
|
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=False)
|
||||||
|
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from typing import Any, Mapping, Optional
|
from typing import Any, Dict, Mapping, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -59,7 +59,9 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
else config.dim // config.n_heads
|
else config.dim // config.n_heads
|
||||||
)
|
)
|
||||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||||
self.rotary_embedding = RotaryEmbedding(rope_dim, config.max_len, rope_base)
|
self.rotary_embedding = RotaryEmbedding(
|
||||||
|
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||||
|
)
|
||||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||||
|
|
||||||
self.layers = nn.ModuleList(
|
self.layers = nn.ModuleList(
|
||||||
@@ -134,7 +136,7 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
input_mask: Optional[Tensor] = None,
|
input_mask: Optional[Tensor] = None,
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
paged_cache: Optional[KvcacheView] = None,
|
||||||
position_ids: Optional[Tensor] = None,
|
position_ids: Optional[Tensor] = None,
|
||||||
) -> Tensor:
|
) -> Dict[str, Tensor]:
|
||||||
assert input_ids.ndim == 2
|
assert input_ids.ndim == 2
|
||||||
|
|
||||||
x = self.embed_tokens(input_ids)
|
x = self.embed_tokens(input_ids)
|
||||||
|
|||||||
@@ -2,7 +2,9 @@ from astrai.parallel.executor import (
|
|||||||
AccumOptimizer,
|
AccumOptimizer,
|
||||||
AccumScheduler,
|
AccumScheduler,
|
||||||
BaseExecutor,
|
BaseExecutor,
|
||||||
|
DDPExecutor,
|
||||||
ExecutorFactory,
|
ExecutorFactory,
|
||||||
|
FSDPExecutor,
|
||||||
GradientState,
|
GradientState,
|
||||||
NoneExecutor,
|
NoneExecutor,
|
||||||
)
|
)
|
||||||
@@ -31,4 +33,6 @@ __all__ = [
|
|||||||
"AccumOptimizer",
|
"AccumOptimizer",
|
||||||
"AccumScheduler",
|
"AccumScheduler",
|
||||||
"NoneExecutor",
|
"NoneExecutor",
|
||||||
|
"DDPExecutor",
|
||||||
|
"FSDPExecutor",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from typing import Optional, Tuple
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||||
from torch.optim import Optimizer
|
from torch.optim import Optimizer
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
@@ -198,3 +199,69 @@ class DDPExecutor(BaseExecutor):
|
|||||||
if isinstance(model, DDP):
|
if isinstance(model, DDP):
|
||||||
return model.module
|
return model.module
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
@ExecutorFactory.register("fsdp")
|
||||||
|
class FSDPExecutor(BaseExecutor):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
grad_accum_steps: int = 1,
|
||||||
|
process_group=None,
|
||||||
|
sharding_strategy=None,
|
||||||
|
cpu_offload=None,
|
||||||
|
auto_wrap_policy=None,
|
||||||
|
backward_prefetch=None,
|
||||||
|
mixed_precision=None,
|
||||||
|
ignored_modules=None,
|
||||||
|
param_init_fn=None,
|
||||||
|
sync_module_states: bool = False,
|
||||||
|
forward_prefetch: bool = False,
|
||||||
|
limit_all_gathers: bool = True,
|
||||||
|
use_orig_params: bool = False,
|
||||||
|
ignored_states=None,
|
||||||
|
device_mesh=None,
|
||||||
|
):
|
||||||
|
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||||
|
self._fsdp_kwargs = {
|
||||||
|
k: v
|
||||||
|
for k, v in dict(
|
||||||
|
process_group=process_group,
|
||||||
|
sharding_strategy=sharding_strategy,
|
||||||
|
cpu_offload=cpu_offload,
|
||||||
|
auto_wrap_policy=auto_wrap_policy,
|
||||||
|
backward_prefetch=backward_prefetch,
|
||||||
|
mixed_precision=mixed_precision,
|
||||||
|
ignored_modules=ignored_modules,
|
||||||
|
param_init_fn=param_init_fn,
|
||||||
|
sync_module_states=sync_module_states,
|
||||||
|
forward_prefetch=forward_prefetch,
|
||||||
|
limit_all_gathers=limit_all_gathers,
|
||||||
|
use_orig_params=use_orig_params,
|
||||||
|
ignored_states=ignored_states,
|
||||||
|
device_mesh=device_mesh,
|
||||||
|
).items()
|
||||||
|
if v is not None
|
||||||
|
}
|
||||||
|
self._original_model: Optional[nn.Module] = None
|
||||||
|
|
||||||
|
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||||
|
if not self.use_distributed:
|
||||||
|
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
||||||
|
return model
|
||||||
|
self._original_model = model
|
||||||
|
device_id = torch.device("cuda", get_rank())
|
||||||
|
model = FSDP(model, device_id=device_id, **self._fsdp_kwargs)
|
||||||
|
logger.info("Model wrapped with FSDP (world_size=%d)", get_world_size())
|
||||||
|
return model
|
||||||
|
|
||||||
|
def _no_sync(self, model: nn.Module):
|
||||||
|
if isinstance(model, FSDP):
|
||||||
|
return model.no_sync()
|
||||||
|
return contextlib.nullcontext()
|
||||||
|
|
||||||
|
def unwrap_model(self, model: nn.Module) -> nn.Module:
|
||||||
|
if self._original_model is not None:
|
||||||
|
return self._original_model
|
||||||
|
if isinstance(model, FSDP):
|
||||||
|
return model._fsdp_wrapped_module
|
||||||
|
return model
|
||||||
|
|||||||
+150
-51
@@ -1,7 +1,9 @@
|
|||||||
|
import io
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Union
|
||||||
|
|
||||||
import safetensors.torch as st
|
import safetensors.torch as st
|
||||||
import torch
|
import torch
|
||||||
@@ -9,75 +11,172 @@ import torch.distributed as dist
|
|||||||
|
|
||||||
from astrai.parallel.setup import get_rank
|
from astrai.parallel.setup import get_rank
|
||||||
|
|
||||||
|
_META_FILE = "meta.json"
|
||||||
|
_CONFIG_FILE = "config.json"
|
||||||
|
_WEIGHTS_FILE = "model.safetensors"
|
||||||
|
|
||||||
class Checkpoint:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
state_dict: Dict[str, Any],
|
|
||||||
epoch: int = 0,
|
|
||||||
iteration: int = 0,
|
|
||||||
extra: Optional[Dict[str, Any]] = None,
|
|
||||||
meta: Optional[Dict[str, Any]] = None,
|
|
||||||
):
|
|
||||||
self.state_dict = state_dict
|
|
||||||
self.epoch = epoch
|
|
||||||
self.iteration = iteration
|
|
||||||
self.extra = extra or {}
|
|
||||||
self.meta = meta or {}
|
|
||||||
|
|
||||||
def save(
|
def save_safetensors(state_dict: dict, path: Union[str, Path]):
|
||||||
self,
|
st.save_file(state_dict, str(path))
|
||||||
save_dir: str,
|
|
||||||
) -> None:
|
|
||||||
|
|
||||||
save_path = Path(save_dir)
|
|
||||||
save_path.mkdir(parents=True, exist_ok=True)
|
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
if not broadcast or not dist.is_initialized():
|
||||||
|
return st.load_file(str(path))
|
||||||
|
|
||||||
rank = get_rank()
|
rank = get_rank()
|
||||||
if rank == 0:
|
if rank == 0:
|
||||||
|
state_dict = st.load_file(str(path))
|
||||||
|
else:
|
||||||
|
state_dict = {}
|
||||||
|
tmp = [state_dict]
|
||||||
|
dist.broadcast_object_list(tmp, src=0)
|
||||||
|
return tmp[0]
|
||||||
|
|
||||||
|
|
||||||
|
def save_json(data: dict, path: Union[str, Path]):
|
||||||
|
with open(str(path), "w") as f:
|
||||||
|
json.dump(data, f, indent=2)
|
||||||
|
|
||||||
|
|
||||||
|
def load_json(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
if not broadcast or not dist.is_initialized():
|
||||||
|
with open(str(path), "r") as f:
|
||||||
|
return json.load(f)
|
||||||
|
|
||||||
|
rank = get_rank()
|
||||||
|
if rank == 0:
|
||||||
|
with open(str(path), "r") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
else:
|
||||||
|
data = {}
|
||||||
|
tmp = [data]
|
||||||
|
dist.broadcast_object_list(tmp, src=0)
|
||||||
|
return tmp[0]
|
||||||
|
|
||||||
|
|
||||||
|
def save_torch(obj: Any, path: Union[str, Path]):
|
||||||
|
torch.save(obj, str(path))
|
||||||
|
|
||||||
|
|
||||||
|
def load_torch(path: Union[str, Path], broadcast: bool = False) -> Any:
|
||||||
|
if not broadcast or not dist.is_initialized():
|
||||||
|
return torch.load(str(path), map_location="cpu", weights_only=False)
|
||||||
|
|
||||||
|
path = Path(path)
|
||||||
|
rank = get_rank()
|
||||||
|
|
||||||
|
if rank == 0:
|
||||||
|
with open(path, "rb") as f:
|
||||||
|
raw = f.read()
|
||||||
|
data_tensor = torch.frombuffer(bytearray(raw), dtype=torch.uint8)
|
||||||
|
num_bytes = torch.tensor([len(raw)], dtype=torch.long)
|
||||||
|
else:
|
||||||
|
num_bytes = torch.tensor([0], dtype=torch.long)
|
||||||
|
|
||||||
|
dist.broadcast(num_bytes, src=0)
|
||||||
|
|
||||||
|
if rank != 0:
|
||||||
|
data_tensor = torch.empty(num_bytes.item(), dtype=torch.uint8)
|
||||||
|
|
||||||
|
dist.broadcast(data_tensor, src=0)
|
||||||
|
|
||||||
|
buf = io.BytesIO(data_tensor.numpy().tobytes())
|
||||||
|
return torch.load(buf, map_location="cpu", weights_only=False)
|
||||||
|
|
||||||
|
|
||||||
|
def save_model(config: dict, state_dict: dict, save_directory: str):
|
||||||
|
save_path = Path(save_directory)
|
||||||
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
save_json(config, save_path / _CONFIG_FILE)
|
||||||
|
save_safetensors(state_dict, save_path / _WEIGHTS_FILE)
|
||||||
|
|
||||||
|
|
||||||
|
def load_model_config(save_directory: str) -> dict:
|
||||||
|
return load_json(Path(save_directory) / _CONFIG_FILE)
|
||||||
|
|
||||||
|
|
||||||
|
def load_model_weights(save_directory: str) -> dict:
|
||||||
|
return load_state_dict(Path(save_directory) / _WEIGHTS_FILE)
|
||||||
|
|
||||||
|
|
||||||
|
def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
path = Path(path)
|
||||||
|
if not broadcast or not dist.is_initialized():
|
||||||
|
return load_safetensors(path)
|
||||||
|
|
||||||
|
rank = get_rank()
|
||||||
|
if rank == 0:
|
||||||
|
state_dict = load_safetensors(path)
|
||||||
|
specs = [
|
||||||
|
(k, list(state_dict[k].shape), str(state_dict[k].dtype).split(".")[-1])
|
||||||
|
for k in sorted(state_dict)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
state_dict = {}
|
||||||
|
specs = []
|
||||||
|
|
||||||
|
specs_list = [specs]
|
||||||
|
dist.broadcast_object_list(specs_list, src=0)
|
||||||
|
specs = specs_list[0]
|
||||||
|
|
||||||
|
for key, shape, dtype_name in specs:
|
||||||
|
dtype = getattr(torch, dtype_name)
|
||||||
|
if rank != 0:
|
||||||
|
tensor = torch.empty(shape, dtype=dtype, device="cpu")
|
||||||
|
else:
|
||||||
|
tensor = state_dict[key].contiguous().cpu()
|
||||||
|
dist.broadcast(tensor, src=0)
|
||||||
|
if rank != 0:
|
||||||
|
state_dict[key] = tensor
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Checkpoint:
|
||||||
|
state_dict: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
epoch: int = 0
|
||||||
|
iteration: int = 0
|
||||||
|
extra: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
meta: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
config: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def save(self, save_dir: str):
|
||||||
|
save_path = Path(save_dir)
|
||||||
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
if get_rank() != 0:
|
||||||
|
return
|
||||||
|
|
||||||
meta = {
|
meta = {
|
||||||
"epoch": self.epoch,
|
"epoch": self.epoch,
|
||||||
"iteration": self.iteration,
|
"iteration": self.iteration,
|
||||||
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||||
|
**self.meta,
|
||||||
}
|
}
|
||||||
meta.update(self.meta)
|
save_json(meta, save_path / _META_FILE)
|
||||||
with open(save_path / "meta.json", "w") as f:
|
save_json(self.config, save_path / _CONFIG_FILE)
|
||||||
json.dump(meta, f, indent=2)
|
save_safetensors(self.state_dict, save_path / _WEIGHTS_FILE)
|
||||||
|
|
||||||
st.save_file(self.state_dict, save_path / "state_dict.safetensors")
|
|
||||||
if self.extra:
|
|
||||||
for key, value in self.extra.items():
|
for key, value in self.extra.items():
|
||||||
torch.save(value, save_path / f"{key}.pt")
|
save_torch(value, save_path / f"{key}.pt")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def load(
|
def load(cls, save_dir: str, broadcast: bool = False) -> "Checkpoint":
|
||||||
cls,
|
|
||||||
save_dir: str,
|
|
||||||
) -> "Checkpoint":
|
|
||||||
|
|
||||||
rank = get_rank()
|
|
||||||
save_path = Path(save_dir)
|
save_path = Path(save_dir)
|
||||||
|
|
||||||
meta = {}
|
meta = load_json(save_path / _META_FILE, broadcast)
|
||||||
if rank == 0:
|
config = load_json(save_path / _CONFIG_FILE, broadcast)
|
||||||
with open(Path(save_dir) / "meta.json", "r") as f:
|
state_dict = load_state_dict(save_path / _WEIGHTS_FILE, broadcast=broadcast)
|
||||||
meta = json.load(f)
|
|
||||||
|
|
||||||
if dist.is_initialized():
|
|
||||||
meta_list = [meta]
|
|
||||||
dist.broadcast_object_list(meta_list, src=0)
|
|
||||||
meta = meta_list[0]
|
|
||||||
|
|
||||||
state_dict = st.load_file(save_path / "state_dict.safetensors")
|
|
||||||
|
|
||||||
extra = {}
|
extra = {}
|
||||||
for f in save_path.iterdir():
|
for f in sorted(save_path.iterdir()):
|
||||||
if f.suffix == ".pt" and f.stem not in ("meta",):
|
if f.suffix == ".pt":
|
||||||
extra[f.stem] = torch.load(f, map_location="cpu", weights_only=False)
|
extra[f.stem] = load_torch(f, broadcast=broadcast)
|
||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
state_dict=state_dict,
|
state_dict=state_dict,
|
||||||
epoch=meta["epoch"],
|
epoch=meta.get("epoch", 0),
|
||||||
iteration=meta["iteration"],
|
iteration=meta.get("iteration", 0),
|
||||||
extra=extra or None,
|
extra=extra,
|
||||||
|
config=config,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ class SchedulerFactory(BaseFactory["BaseScheduler"]):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _validate_component(cls, scheduler_cls: Type[BaseScheduler]) -> None:
|
def _validate_component(cls, scheduler_cls: Type[BaseScheduler]):
|
||||||
"""Validate that the scheduler class inherits from BaseScheduler."""
|
"""Validate that the scheduler class inherits from BaseScheduler."""
|
||||||
if not issubclass(scheduler_cls, BaseScheduler):
|
if not issubclass(scheduler_cls, BaseScheduler):
|
||||||
raise TypeError(f"{scheduler_cls.__name__} must inherit from BaseScheduler")
|
raise TypeError(f"{scheduler_cls.__name__} must inherit from BaseScheduler")
|
||||||
|
|||||||
@@ -8,15 +8,17 @@ 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.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
def unwrap_model(model: nn.Module) -> nn.Module:
|
def unwrap_model(model: nn.Module) -> nn.Module:
|
||||||
"""Unwrap DDP wrapper if present to get the original model."""
|
|
||||||
if isinstance(model, DDP):
|
if isinstance(model, DDP):
|
||||||
return model.module
|
return model.module
|
||||||
|
if isinstance(model, FSDP):
|
||||||
|
return model._fsdp_wrapped_module
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
@@ -123,7 +125,7 @@ class StrategyFactory(BaseFactory["BaseStrategy"]):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _validate_component(cls, strategy_cls: type) -> None:
|
def _validate_component(cls, strategy_cls: type):
|
||||||
"""Validate that the strategy class inherits from BaseStrategy."""
|
"""Validate that the strategy class inherits from BaseStrategy."""
|
||||||
if not issubclass(strategy_cls, BaseStrategy):
|
if not issubclass(strategy_cls, BaseStrategy):
|
||||||
raise TypeError(f"{strategy_cls.__name__} must inherit from BaseStrategy")
|
raise TypeError(f"{strategy_cls.__name__} must inherit from BaseStrategy")
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from tqdm import tqdm
|
|||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.parallel import only_on_rank
|
from astrai.parallel import only_on_rank
|
||||||
from astrai.parallel.setup import get_current_device
|
from astrai.parallel.setup import get_current_device, get_rank
|
||||||
from astrai.serialization import Checkpoint
|
from astrai.serialization import Checkpoint
|
||||||
from astrai.trainer.metric_util import (
|
from astrai.trainer.metric_util import (
|
||||||
ctx_get_grad_max,
|
ctx_get_grad_max,
|
||||||
@@ -137,44 +137,32 @@ class CheckpointCallback(TrainCallback):
|
|||||||
save_dir: str,
|
save_dir: str,
|
||||||
interval: int,
|
interval: int,
|
||||||
weight_only: bool = False,
|
weight_only: bool = False,
|
||||||
state_dict_fn: Optional[Callable[[nn.Module], dict]] = None,
|
|
||||||
save_extra_fn: Optional[Callable[["TrainContext"], dict]] = None,
|
save_extra_fn: Optional[Callable[["TrainContext"], dict]] = None,
|
||||||
load_extra_fn: Optional[Callable[[dict, "TrainContext"], None]] = None,
|
|
||||||
):
|
):
|
||||||
self.save_dir = save_dir
|
self.save_dir = save_dir
|
||||||
self.interval = interval
|
self.interval = interval
|
||||||
self.weight_only = weight_only
|
self.weight_only = weight_only
|
||||||
self.state_dict_fn = state_dict_fn
|
|
||||||
self.save_extra_fn = save_extra_fn or CheckpointCallback.save_extra
|
self.save_extra_fn = save_extra_fn or CheckpointCallback.save_extra
|
||||||
self.load_extra_fn = load_extra_fn or CheckpointCallback.load_extra
|
|
||||||
self.last_ckpt_iter = 0
|
self.last_ckpt_iter = 0
|
||||||
|
|
||||||
@only_on_rank(0)
|
|
||||||
def _save_checkpoint(self, context: TrainContext):
|
def _save_checkpoint(self, context: TrainContext):
|
||||||
|
unwrapped = context.executor.unwrap_model(context.model)
|
||||||
|
state_dict = unwrapped.state_dict()
|
||||||
|
self.last_ckpt_iter = context.iteration
|
||||||
|
|
||||||
|
if get_rank() == 0:
|
||||||
save_path = os.path.join(
|
save_path = os.path.join(
|
||||||
self.save_dir, f"epoch_{context.epoch}_iter_{context.iteration}"
|
self.save_dir, f"epoch_{context.epoch}_iter_{context.iteration}"
|
||||||
)
|
)
|
||||||
state_dict = (
|
|
||||||
self.state_dict_fn(context.model)
|
|
||||||
if self.state_dict_fn
|
|
||||||
else context.model.state_dict()
|
|
||||||
)
|
|
||||||
|
|
||||||
extra = self.save_extra_fn(context)
|
extra = self.save_extra_fn(context)
|
||||||
context.checkpoint = Checkpoint(
|
context.checkpoint = Checkpoint(
|
||||||
state_dict=state_dict,
|
state_dict=state_dict,
|
||||||
epoch=context.epoch,
|
epoch=context.epoch,
|
||||||
iteration=context.iteration,
|
iteration=context.iteration,
|
||||||
extra=extra,
|
extra=extra,
|
||||||
meta=context.config.to_dict(),
|
config=context.model_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
context.checkpoint.save(save_path)
|
context.checkpoint.save(save_path)
|
||||||
self.last_ckpt_iter = context.iteration
|
|
||||||
|
|
||||||
def on_train_begin(self, context: TrainContext):
|
|
||||||
if context.checkpoint and context.checkpoint.extra:
|
|
||||||
self.load_extra_fn(context.checkpoint.extra, context)
|
|
||||||
|
|
||||||
def on_batch_end(self, context: TrainContext):
|
def on_batch_end(self, context: TrainContext):
|
||||||
if context.iteration - self.last_ckpt_iter >= self.interval:
|
if context.iteration - self.last_ckpt_iter >= self.interval:
|
||||||
@@ -196,12 +184,6 @@ class CheckpointCallback(TrainCallback):
|
|||||||
extra[name] = obj.state_dict()
|
extra[name] = obj.state_dict()
|
||||||
return extra
|
return extra
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def load_extra(extra: dict, context: TrainContext):
|
|
||||||
for name in CheckpointCallback.extra_keys:
|
|
||||||
if name in extra:
|
|
||||||
getattr(context, name).load_state_dict(extra[name])
|
|
||||||
|
|
||||||
|
|
||||||
@CallbackFactory.register("progress_bar")
|
@CallbackFactory.register("progress_bar")
|
||||||
class ProgressBarCallback(TrainCallback):
|
class ProgressBarCallback(TrainCallback):
|
||||||
@@ -210,7 +192,7 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, num_epoch: int, log_interval: int = 100, file: IO[str] = sys.stdout
|
self, num_epoch: int, log_interval: int = 100, file: Optional[IO[str]] = None
|
||||||
):
|
):
|
||||||
self.num_epoch = num_epoch
|
self.num_epoch = num_epoch
|
||||||
self.log_interval = log_interval
|
self.log_interval = log_interval
|
||||||
@@ -223,7 +205,7 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
context.dataloader,
|
context.dataloader,
|
||||||
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
|
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
|
||||||
dynamic_ncols=True,
|
dynamic_ncols=True,
|
||||||
file=self.file,
|
file=self.file or sys.stdout,
|
||||||
)
|
)
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
from pathlib import Path
|
||||||
from typing import Optional, Self
|
from typing import Optional, Self
|
||||||
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -6,10 +7,11 @@ from torch.utils.data import DataLoader
|
|||||||
|
|
||||||
from astrai.config.train_config import TrainConfig
|
from astrai.config.train_config import TrainConfig
|
||||||
from astrai.dataset import ResumableDistributedSampler
|
from astrai.dataset import ResumableDistributedSampler
|
||||||
|
from astrai.model.components.lora import inject_lora
|
||||||
from astrai.parallel.executor import BaseExecutor, ExecutorFactory
|
from astrai.parallel.executor import BaseExecutor, ExecutorFactory
|
||||||
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
||||||
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
||||||
from astrai.serialization import Checkpoint
|
from astrai.serialization import Checkpoint, load_json, load_model_weights
|
||||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||||
|
|
||||||
|
|
||||||
@@ -22,6 +24,7 @@ class TrainContext:
|
|||||||
scheduler: SchedulerProtocol = field(default=None)
|
scheduler: SchedulerProtocol = field(default=None)
|
||||||
checkpoint: Checkpoint = field(default=None)
|
checkpoint: Checkpoint = field(default=None)
|
||||||
config: TrainConfig = field(default=None)
|
config: TrainConfig = field(default=None)
|
||||||
|
model_config: dict = field(default_factory=dict)
|
||||||
executor: BaseExecutor = field(default=None)
|
executor: BaseExecutor = field(default=None)
|
||||||
|
|
||||||
epoch: int = field(default=0)
|
epoch: int = field(default=0)
|
||||||
@@ -41,10 +44,10 @@ class TrainContextBuilder:
|
|||||||
config: TrainConfig,
|
config: TrainConfig,
|
||||||
):
|
):
|
||||||
self.config = config
|
self.config = config
|
||||||
self._checkpoint: Optional[Checkpoint] = None
|
self._resume_dir: Optional[str] = None
|
||||||
|
|
||||||
def with_checkpoint(self, checkpoint: Optional[Checkpoint]) -> Self:
|
def with_resume_dir(self, resume_dir: Optional[str]) -> Self:
|
||||||
self._checkpoint = checkpoint
|
self._resume_dir = resume_dir
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def build(self) -> TrainContext:
|
def build(self) -> TrainContext:
|
||||||
@@ -57,27 +60,52 @@ class TrainContextBuilder:
|
|||||||
**cfg.executor_kwargs,
|
**cfg.executor_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
model = cfg.model_fn()
|
||||||
|
model = model.to(device=device)
|
||||||
|
|
||||||
|
model_config = {}
|
||||||
|
if self._resume_dir:
|
||||||
|
config_path = Path(self._resume_dir) / "config.json"
|
||||||
|
if config_path.exists():
|
||||||
|
model_config = load_json(config_path)
|
||||||
|
|
||||||
|
if not model_config and hasattr(model, "config"):
|
||||||
|
model_config = model.config.to_dict()
|
||||||
|
|
||||||
context = TrainContext(
|
context = TrainContext(
|
||||||
model=cfg.model,
|
model=model,
|
||||||
world_size=get_world_size(),
|
world_size=get_world_size(),
|
||||||
rank=get_rank(),
|
rank=get_rank(),
|
||||||
config=cfg,
|
config=cfg,
|
||||||
|
model_config=model_config,
|
||||||
executor=executor,
|
executor=executor,
|
||||||
)
|
)
|
||||||
|
|
||||||
context.model = context.model.to(device=device)
|
if self._resume_dir is not None:
|
||||||
|
resume_path = Path(self._resume_dir)
|
||||||
if self._checkpoint is not None:
|
if (resume_path / "meta.json").exists():
|
||||||
context.epoch = max(self._checkpoint.epoch, cfg.start_epoch)
|
checkpoint = Checkpoint.load(self._resume_dir)
|
||||||
context.iteration = max(self._checkpoint.iteration, cfg.start_batch)
|
state_dict = checkpoint.state_dict
|
||||||
context.model.load_state_dict(self._checkpoint.state_dict)
|
if checkpoint.config:
|
||||||
context.checkpoint = self._checkpoint
|
context.model_config = checkpoint.config
|
||||||
else:
|
else:
|
||||||
context.checkpoint = Checkpoint(
|
checkpoint = None
|
||||||
state_dict=context.model.state_dict(),
|
state_dict = load_model_weights(self._resume_dir)
|
||||||
|
model.load_state_dict(state_dict, strict=False)
|
||||||
|
if checkpoint is not None:
|
||||||
|
context.epoch = cfg.start_epoch
|
||||||
|
context.iteration = cfg.start_batch
|
||||||
|
context.checkpoint = checkpoint
|
||||||
|
|
||||||
|
if cfg.lora is not None:
|
||||||
|
inject_lora(
|
||||||
|
model,
|
||||||
|
r=cfg.lora.r,
|
||||||
|
alpha=cfg.lora.alpha,
|
||||||
|
target_modules=set(cfg.lora.target_modules),
|
||||||
)
|
)
|
||||||
|
|
||||||
context.optimizer = cfg.optimizer_fn(context.model)
|
context.optimizer = cfg.optimizer_fn(model)
|
||||||
context.scheduler = cfg.scheduler_fn(context.optimizer)
|
context.scheduler = cfg.scheduler_fn(context.optimizer)
|
||||||
|
|
||||||
sampler_offset = context.iteration * cfg.batch_per_device
|
sampler_offset = context.iteration * cfg.batch_per_device
|
||||||
@@ -115,13 +143,21 @@ class TrainContextBuilder:
|
|||||||
|
|
||||||
context.model, context.optimizer, context.dataloader, context.scheduler = (
|
context.model, context.optimizer, context.dataloader, context.scheduler = (
|
||||||
executor.prepare(
|
executor.prepare(
|
||||||
context.model,
|
model,
|
||||||
context.optimizer,
|
context.optimizer,
|
||||||
context.dataloader,
|
context.dataloader,
|
||||||
context.scheduler,
|
context.scheduler,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if context.checkpoint and context.checkpoint.extra:
|
||||||
|
extra = context.checkpoint.extra
|
||||||
|
for name in ("optimizer", "scheduler"):
|
||||||
|
if name in extra:
|
||||||
|
obj = getattr(context, name, None)
|
||||||
|
if obj is not None:
|
||||||
|
obj.load_state_dict(extra[name])
|
||||||
|
|
||||||
context.strategy = StrategyFactory.create(
|
context.strategy = StrategyFactory.create(
|
||||||
model=context.model,
|
model=context.model,
|
||||||
train_type=cfg.strategy,
|
train_type=cfg.strategy,
|
||||||
|
|||||||
@@ -3,7 +3,6 @@ from typing import List, Optional
|
|||||||
|
|
||||||
from astrai.config import TrainConfig
|
from astrai.config import TrainConfig
|
||||||
from astrai.parallel.setup import spawn_parallel_fn
|
from astrai.parallel.setup import spawn_parallel_fn
|
||||||
from astrai.serialization import Checkpoint
|
|
||||||
from astrai.trainer.train_callback import (
|
from astrai.trainer.train_callback import (
|
||||||
CallbackFactory,
|
CallbackFactory,
|
||||||
TrainCallback,
|
TrainCallback,
|
||||||
@@ -54,9 +53,9 @@ class Trainer:
|
|||||||
if method:
|
if method:
|
||||||
method(context)
|
method(context)
|
||||||
|
|
||||||
def _trainer_loop(self, checkpoint: Optional[Checkpoint] = None):
|
def _trainer_loop(self, resume_dir: Optional[str] = None):
|
||||||
context = (
|
context = (
|
||||||
TrainContextBuilder(self.train_config).with_checkpoint(checkpoint).build()
|
TrainContextBuilder(self.train_config).with_resume_dir(resume_dir).build()
|
||||||
)
|
)
|
||||||
executor = context.executor
|
executor = context.executor
|
||||||
self._call_callbacks("on_train_begin", context)
|
self._call_callbacks("on_train_begin", context)
|
||||||
@@ -90,13 +89,13 @@ class Trainer:
|
|||||||
self._call_callbacks("on_epoch_end", context)
|
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("Training failed: %s", str(e), exc_info=True)
|
||||||
self._call_callbacks("on_error", context)
|
self._call_callbacks("on_error", context)
|
||||||
raise
|
raise
|
||||||
finally:
|
finally:
|
||||||
self._call_callbacks("on_train_end", context)
|
self._call_callbacks("on_train_end", context)
|
||||||
|
|
||||||
def train(self, checkpoint: Optional[Checkpoint] = None):
|
def train(self, resume_dir: Optional[str] = None):
|
||||||
cfg = self.train_config
|
cfg = self.train_config
|
||||||
spawn_parallel_fn(
|
spawn_parallel_fn(
|
||||||
self._trainer_loop,
|
self._trainer_loop,
|
||||||
@@ -106,5 +105,5 @@ class Trainer:
|
|||||||
master_port=cfg.master_port,
|
master_port=cfg.master_port,
|
||||||
device_type=cfg.device_type,
|
device_type=cfg.device_type,
|
||||||
start_method=cfg.start_method,
|
start_method=cfg.start_method,
|
||||||
checkpoint=checkpoint,
|
resume_dir=resume_dir,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,279 @@
|
|||||||
|
"""MMLU evaluation via log-likelihood ranking."""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import csv
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
import urllib.request
|
||||||
|
import zipfile
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
import tqdm
|
||||||
|
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
MMLU_URL = "https://github.com/hendrycks/test/archive/refs/heads/master.zip"
|
||||||
|
MMLU_SUBJECTS = [
|
||||||
|
"abstract_algebra",
|
||||||
|
"anatomy",
|
||||||
|
"astronomy",
|
||||||
|
"business_ethics",
|
||||||
|
"clinical_knowledge",
|
||||||
|
"college_biology",
|
||||||
|
"college_chemistry",
|
||||||
|
"college_computer_science",
|
||||||
|
"college_mathematics",
|
||||||
|
"college_medicine",
|
||||||
|
"college_physics",
|
||||||
|
"computer_security",
|
||||||
|
"conceptual_physics",
|
||||||
|
"econometrics",
|
||||||
|
"electrical_engineering",
|
||||||
|
"elementary_mathematics",
|
||||||
|
"formal_logic",
|
||||||
|
"global_facts",
|
||||||
|
"high_school_biology",
|
||||||
|
"high_school_chemistry",
|
||||||
|
"high_school_computer_science",
|
||||||
|
"high_school_european_history",
|
||||||
|
"high_school_geography",
|
||||||
|
"high_school_government_and_politics",
|
||||||
|
"high_school_macroeconomics",
|
||||||
|
"high_school_mathematics",
|
||||||
|
"high_school_microeconomics",
|
||||||
|
"high_school_physics",
|
||||||
|
"high_school_psychology",
|
||||||
|
"high_school_statistics",
|
||||||
|
"high_school_us_history",
|
||||||
|
"high_school_world_history",
|
||||||
|
"human_aging",
|
||||||
|
"human_sexuality",
|
||||||
|
"international_law",
|
||||||
|
"jurisprudence",
|
||||||
|
"logical_fallacies",
|
||||||
|
"machine_learning",
|
||||||
|
"management",
|
||||||
|
"marketing",
|
||||||
|
"medical_genetics",
|
||||||
|
"miscellaneous",
|
||||||
|
"moral_disputes",
|
||||||
|
"moral_scenarios",
|
||||||
|
"nutrition",
|
||||||
|
"philosophy",
|
||||||
|
"prehistory",
|
||||||
|
"professional_accounting",
|
||||||
|
"professional_law",
|
||||||
|
"professional_medicine",
|
||||||
|
"professional_psychology",
|
||||||
|
"public_relations",
|
||||||
|
"security_studies",
|
||||||
|
"sociology",
|
||||||
|
"us_foreign_policy",
|
||||||
|
"virology",
|
||||||
|
"world_religions",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _download_and_extract(url: str, data_dir: str):
|
||||||
|
zip_path = os.path.join(data_dir, "mmlu.zip")
|
||||||
|
os.makedirs(data_dir, exist_ok=True)
|
||||||
|
print(f"Downloading MMLU data from {url}...")
|
||||||
|
urllib.request.urlretrieve(url, zip_path)
|
||||||
|
print("Extracting...")
|
||||||
|
with zipfile.ZipFile(zip_path, "r") as zf:
|
||||||
|
zf.extractall(data_dir)
|
||||||
|
os.remove(zip_path)
|
||||||
|
|
||||||
|
|
||||||
|
def download_mmlu(data_dir: str):
|
||||||
|
_download_and_extract(MMLU_URL, data_dir)
|
||||||
|
src = os.path.join(data_dir, "test-master", "data")
|
||||||
|
if os.path.exists(src):
|
||||||
|
for item in os.listdir(src):
|
||||||
|
os.rename(os.path.join(src, item), os.path.join(data_dir, item))
|
||||||
|
shutil.rmtree(os.path.join(data_dir, "test-master"))
|
||||||
|
print(f"MMLU data saved to {data_dir}")
|
||||||
|
|
||||||
|
|
||||||
|
def _strip_prefix(text: str, prefix: str) -> str:
|
||||||
|
if text.startswith(prefix):
|
||||||
|
return text[len(prefix) :].strip()
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def load_csv(path: str) -> list[dict]:
|
||||||
|
data = []
|
||||||
|
with open(path, "r", encoding="utf-8") as f:
|
||||||
|
for row in csv.reader(f):
|
||||||
|
if len(row) < 6:
|
||||||
|
continue
|
||||||
|
if row[0].strip().lower() == "question":
|
||||||
|
continue
|
||||||
|
data.append(
|
||||||
|
{
|
||||||
|
"question": row[0].strip(),
|
||||||
|
"A": _strip_prefix(row[1].strip(), "A)"),
|
||||||
|
"B": _strip_prefix(row[2].strip(), "B)"),
|
||||||
|
"C": _strip_prefix(row[3].strip(), "C)"),
|
||||||
|
"D": _strip_prefix(row[4].strip(), "D)"),
|
||||||
|
"answer": row[5].strip(),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
def build_prompt(
|
||||||
|
question: str, choices: dict, subject: str, n_shot: int, dev_data: list[dict]
|
||||||
|
) -> str:
|
||||||
|
prompt = ""
|
||||||
|
if n_shot > 0 and dev_data:
|
||||||
|
prompt = f"The following are multiple choice questions (with answers) about {subject}.\n\n"
|
||||||
|
for item in dev_data[:n_shot]:
|
||||||
|
prompt += f"Question: {item['question']}\n"
|
||||||
|
for k in ("A", "B", "C", "D"):
|
||||||
|
prompt += f"{k}. {item[k]}\n"
|
||||||
|
prompt += f"Answer: {item['answer']}\n\n"
|
||||||
|
prompt += f"Question: {question}\n"
|
||||||
|
for k in ("A", "B", "C", "D"):
|
||||||
|
prompt += f"{k}. {choices[k]}\n"
|
||||||
|
prompt += "Answer:"
|
||||||
|
return prompt
|
||||||
|
|
||||||
|
|
||||||
|
def choice_logprob(
|
||||||
|
model, tokenizer, context_ids: list[int], choice_letter: str, device: str
|
||||||
|
) -> float:
|
||||||
|
choice_text = f" {choice_letter}"
|
||||||
|
choice_ids = tokenizer.encode(choice_text, add_special_tokens=False)
|
||||||
|
input_ids = context_ids + choice_ids
|
||||||
|
max_len = model.config.max_len
|
||||||
|
if len(input_ids) > max_len:
|
||||||
|
overflow = len(input_ids) - max_len
|
||||||
|
input_ids = input_ids[overflow:]
|
||||||
|
ctx_len = len(input_ids) - len(choice_ids)
|
||||||
|
else:
|
||||||
|
ctx_len = len(context_ids)
|
||||||
|
|
||||||
|
input_tensor = torch.tensor([input_ids], device=device, dtype=torch.long)
|
||||||
|
with torch.inference_mode():
|
||||||
|
logits = model(input_tensor)["logits"][0]
|
||||||
|
|
||||||
|
score = 0.0
|
||||||
|
for i, tid in enumerate(choice_ids):
|
||||||
|
pos = ctx_len - 1 + i
|
||||||
|
if pos >= len(logits):
|
||||||
|
break
|
||||||
|
score += F.log_softmax(logits[pos], dim=-1)[tid].item()
|
||||||
|
return score
|
||||||
|
|
||||||
|
|
||||||
|
def evaluate_subject(
|
||||||
|
model,
|
||||||
|
tokenizer,
|
||||||
|
subject: str,
|
||||||
|
test_data: list[dict],
|
||||||
|
dev_data: list[dict] | None,
|
||||||
|
device: str,
|
||||||
|
n_shot: int,
|
||||||
|
) -> tuple[float, int, int]:
|
||||||
|
correct = 0
|
||||||
|
total = 0
|
||||||
|
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
|
||||||
|
prompt = build_prompt(item["question"], item, subject, n_shot, dev_data or [])
|
||||||
|
context_ids = tokenizer.encode(prompt)
|
||||||
|
scores = {
|
||||||
|
c: choice_logprob(model, tokenizer, context_ids, c, device)
|
||||||
|
for c in ("A", "B", "C", "D")
|
||||||
|
}
|
||||||
|
if max(scores, key=scores.get) == item["answer"]:
|
||||||
|
correct += 1
|
||||||
|
total += 1
|
||||||
|
return correct / total, correct, total
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser(description="MMLU evaluation")
|
||||||
|
parser.add_argument(
|
||||||
|
"--param_path", type=str, default="./params", help="Model directory"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--data_dir", type=str, default="./mmlu_data", help="MMLU data directory"
|
||||||
|
)
|
||||||
|
parser.add_argument("--download", action="store_true", help="Download MMLU data")
|
||||||
|
parser.add_argument(
|
||||||
|
"--n_shot", type=int, default=5, help="Few-shot examples (0 for zero-shot)"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--subjects", type=str, nargs="+", help="Specific subjects (default: all)"
|
||||||
|
)
|
||||||
|
parser.add_argument("--output", type=str, help="Output JSON path")
|
||||||
|
parser.add_argument("--split", type=str, default="test", choices=["test", "val"])
|
||||||
|
parser.add_argument(
|
||||||
|
"--device",
|
||||||
|
type=str,
|
||||||
|
default="cuda" if torch.cuda.is_available() else "cpu",
|
||||||
|
help="Device",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--dtype",
|
||||||
|
type=str,
|
||||||
|
default="bfloat16" if torch.cuda.is_available() else "float32",
|
||||||
|
help="Torch dtype",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
if args.download or not os.path.exists(args.data_dir):
|
||||||
|
download_mmlu(args.data_dir)
|
||||||
|
|
||||||
|
model = AutoModel.from_pretrained(args.param_path)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(args.param_path)
|
||||||
|
device = args.device
|
||||||
|
dtype = getattr(torch, args.dtype)
|
||||||
|
model.to(device=device, dtype=dtype)
|
||||||
|
|
||||||
|
subjects = args.subjects or MMLU_SUBJECTS
|
||||||
|
results = {}
|
||||||
|
total_correct = 0
|
||||||
|
total_questions = 0
|
||||||
|
|
||||||
|
for subject in subjects:
|
||||||
|
dev_path = os.path.join(args.data_dir, "dev", f"{subject}_dev.csv")
|
||||||
|
test_path = os.path.join(
|
||||||
|
args.data_dir, args.split, f"{subject}_{args.split}.csv"
|
||||||
|
)
|
||||||
|
|
||||||
|
if not os.path.exists(test_path):
|
||||||
|
print(f" Skipping {subject}: test file not found")
|
||||||
|
continue
|
||||||
|
|
||||||
|
dev_data = load_csv(dev_path) if os.path.exists(dev_path) else None
|
||||||
|
test_data = load_csv(test_path)
|
||||||
|
|
||||||
|
acc, corr, tot = evaluate_subject(
|
||||||
|
model, tokenizer, subject, test_data, dev_data, device, args.n_shot
|
||||||
|
)
|
||||||
|
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
||||||
|
total_correct += corr
|
||||||
|
total_questions += tot
|
||||||
|
print(f" {subject:40s} {acc:.2%} ({corr}/{tot})")
|
||||||
|
|
||||||
|
overall = total_correct / total_questions if total_questions else 0
|
||||||
|
print(f"\n{'=' * 70}")
|
||||||
|
print(f" Overall: {overall:.2%} ({total_correct}/{total_questions})")
|
||||||
|
results["_overall"] = {
|
||||||
|
"accuracy": round(overall, 4),
|
||||||
|
"correct": total_correct,
|
||||||
|
"total": total_questions,
|
||||||
|
}
|
||||||
|
|
||||||
|
if args.output:
|
||||||
|
with open(args.output, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(results, f, indent=2)
|
||||||
|
print(f"Results saved to {args.output}")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -10,11 +10,11 @@ from astrai.tokenize import AutoTokenizer
|
|||||||
|
|
||||||
|
|
||||||
def process_file(
|
def process_file(
|
||||||
model_dir: str, input_file: str, output_file: str, batch_size: int, text_key: str
|
param_path: str, input_file: str, output_file: str, batch_size: int, text_key: str
|
||||||
):
|
):
|
||||||
# Load model and tokenizer
|
# Load model and tokenizer
|
||||||
model = AutoModel.from_pretrained(model_dir)
|
model = AutoModel.from_pretrained(param_path)
|
||||||
tokenizer = AutoTokenizer.from_pretrained(model_dir)
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
model.to(device="cuda", dtype=torch.bfloat16)
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
|
|
||||||
with open(input_file, "r", encoding="utf-8") as f:
|
with open(input_file, "r", encoding="utf-8") as f:
|
||||||
@@ -44,8 +44,8 @@ def process_file(
|
|||||||
|
|
||||||
for seq in batch_encoded:
|
for seq in batch_encoded:
|
||||||
pad_len = max_len - len(seq)
|
pad_len = max_len - len(seq)
|
||||||
padded_seq = [tokenizer.pad_id] * pad_len + seq
|
padded_seq = seq + [tokenizer.pad_id] * pad_len
|
||||||
mask = [False] * pad_len + [True] * len(seq)
|
mask = [True] * len(seq) + [False] * pad_len
|
||||||
padded_ids.append(padded_seq)
|
padded_ids.append(padded_seq)
|
||||||
masks.append(mask)
|
masks.append(mask)
|
||||||
|
|
||||||
@@ -88,7 +88,7 @@ def process_file(
|
|||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
parser = argparse.ArgumentParser(description="Run perplexity with a Khaosz model.")
|
parser = argparse.ArgumentParser(description="Run perplexity with a Khaosz model.")
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--model_dir", type=str, required=True, help="Path to the model directory."
|
"--param_path", type=str, required=True, help="Path to the model directory."
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--input_file", type=str, required=True, help="Path to the input file."
|
"--input_file", type=str, required=True, help="Path to the input file."
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ def main():
|
|||||||
"--reload", action="store_true", help="Enable auto-reload for development"
|
"--reload", action="store_true", help="Enable auto-reload for development"
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--param-path",
|
"--param_path",
|
||||||
type=Path,
|
type=Path,
|
||||||
default=None,
|
default=None,
|
||||||
help="Path to model parameters (default: project_root/params)",
|
help="Path to model parameters (default: project_root/params)",
|
||||||
|
|||||||
+11
-18
@@ -2,7 +2,6 @@ import argparse
|
|||||||
import os
|
import os
|
||||||
from functools import partial
|
from functools import partial
|
||||||
|
|
||||||
import safetensors.torch as st
|
|
||||||
import torch
|
import torch
|
||||||
import torch.optim as optim
|
import torch.optim as optim
|
||||||
|
|
||||||
@@ -147,8 +146,8 @@ def parse_args() -> argparse.Namespace:
|
|||||||
"--parallel_mode",
|
"--parallel_mode",
|
||||||
type=str,
|
type=str,
|
||||||
default="none",
|
default="none",
|
||||||
choices=["none", "ddp"],
|
choices=["none", "ddp", "fsdp"],
|
||||||
help="Parallel training strategy.",
|
help="Parallel training strategy (none, ddp, fsdp).",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--device_type", type=str, default="cuda", help="Device type to use."
|
"--device_type", type=str, default="cuda", help="Device type to use."
|
||||||
@@ -166,6 +165,10 @@ def parse_args() -> argparse.Namespace:
|
|||||||
return args
|
return args
|
||||||
|
|
||||||
|
|
||||||
|
def create_model(config):
|
||||||
|
return AutoRegressiveLM(config).to(dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
|
||||||
def create_optimizer(model, **kwargs) -> optim.Optimizer:
|
def create_optimizer(model, **kwargs) -> optim.Optimizer:
|
||||||
return optim.AdamW(model.parameters(), fused=True, **kwargs)
|
return optim.AdamW(model.parameters(), fused=True, **kwargs)
|
||||||
|
|
||||||
@@ -228,6 +231,8 @@ def train(
|
|||||||
):
|
):
|
||||||
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
||||||
assert os.path.exists(param_path)
|
assert os.path.exists(param_path)
|
||||||
|
if nprocs > 1 and parallel_mode == "none":
|
||||||
|
raise ValueError("--nprocs > 1 requires --parallel_mode to be 'ddp' or 'fsdp'")
|
||||||
|
|
||||||
# Load config
|
# Load config
|
||||||
config_path = os.path.join(param_path, "config.json")
|
config_path = os.path.join(param_path, "config.json")
|
||||||
@@ -236,17 +241,6 @@ def train(
|
|||||||
if window_size is None:
|
if window_size is None:
|
||||||
window_size = config.max_len
|
window_size = config.max_len
|
||||||
|
|
||||||
# Create bare AutoRegressiveLM (for training, no tokenizer needed)
|
|
||||||
model = AutoRegressiveLM(config)
|
|
||||||
|
|
||||||
# Load weights if available
|
|
||||||
weights_path = os.path.join(param_path, "model.safetensors")
|
|
||||||
if os.path.exists(weights_path):
|
|
||||||
state_dict = st.load_file(weights_path)
|
|
||||||
model.load_state_dict(state_dict, strict=False)
|
|
||||||
|
|
||||||
model = model.to(dtype=torch.bfloat16)
|
|
||||||
|
|
||||||
strategy_kwargs = {
|
strategy_kwargs = {
|
||||||
"beta": dpo_beta,
|
"beta": dpo_beta,
|
||||||
"label_smoothing": label_smoothing,
|
"label_smoothing": label_smoothing,
|
||||||
@@ -257,12 +251,11 @@ def train(
|
|||||||
}
|
}
|
||||||
|
|
||||||
executor_kwargs = {
|
executor_kwargs = {
|
||||||
"static_graph": True,
|
|
||||||
"find_unused_parameters": False,
|
|
||||||
"gradient_as_bucket_view": True,
|
"gradient_as_bucket_view": True,
|
||||||
"broadcast_buffers": False,
|
"broadcast_buffers": False,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
model_fn = partial(create_model, config)
|
||||||
dataset = DatasetFactory.load(
|
dataset = DatasetFactory.load(
|
||||||
train_type=train_type,
|
train_type=train_type,
|
||||||
load_path=data_root_path,
|
load_path=data_root_path,
|
||||||
@@ -294,7 +287,7 @@ def train(
|
|||||||
)
|
)
|
||||||
|
|
||||||
train_config = TrainConfig(
|
train_config = TrainConfig(
|
||||||
model=model,
|
model_fn=model_fn,
|
||||||
strategy=train_type,
|
strategy=train_type,
|
||||||
dataset=dataset,
|
dataset=dataset,
|
||||||
optimizer_fn=optimizer_fn,
|
optimizer_fn=optimizer_fn,
|
||||||
@@ -319,7 +312,7 @@ def train(
|
|||||||
)
|
)
|
||||||
|
|
||||||
trainer = Trainer(train_config)
|
trainer = Trainer(train_config)
|
||||||
trainer.train()
|
trainer.train(resume_dir=param_path)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import os
|
||||||
import tempfile
|
import tempfile
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -36,7 +37,6 @@ def test_single_process():
|
|||||||
|
|
||||||
|
|
||||||
def test_checkpoint_with_extra():
|
def test_checkpoint_with_extra():
|
||||||
"""Verify extra keys are saved as individual .pt files and loaded back."""
|
|
||||||
model = torch.nn.Linear(10, 5)
|
model = torch.nn.Linear(10, 5)
|
||||||
optimizer = AdamW(model.parameters(), lr=1e-3)
|
optimizer = AdamW(model.parameters(), lr=1e-3)
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
@@ -52,8 +52,6 @@ def test_checkpoint_with_extra():
|
|||||||
with tempfile.TemporaryDirectory() as tmpdir:
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
checkpoint.save(tmpdir)
|
checkpoint.save(tmpdir)
|
||||||
|
|
||||||
import os
|
|
||||||
|
|
||||||
assert os.path.exists(os.path.join(tmpdir, "optimizer.pt"))
|
assert os.path.exists(os.path.join(tmpdir, "optimizer.pt"))
|
||||||
assert os.path.exists(os.path.join(tmpdir, "scheduler.pt"))
|
assert os.path.exists(os.path.join(tmpdir, "scheduler.pt"))
|
||||||
|
|
||||||
|
|||||||
+194
-169
@@ -7,12 +7,12 @@ import torch
|
|||||||
|
|
||||||
from astrai.dataset.dataset import DatasetFactory, SEQDataset
|
from astrai.dataset.dataset import DatasetFactory, SEQDataset
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
BaseSegmentFetcher,
|
H5Store,
|
||||||
H5Storage,
|
MmapStore,
|
||||||
MultiSegmentFetcher,
|
StoreFactory,
|
||||||
StorageFactory,
|
|
||||||
detect_format,
|
detect_format,
|
||||||
load_json,
|
load_bin,
|
||||||
|
save_bin,
|
||||||
save_h5,
|
save_h5,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -157,111 +157,6 @@ def test_dataset_with_custom_stride(base_test_env):
|
|||||||
assert len(dataset) > len(default_stride_dataset)
|
assert len(dataset) > len(default_stride_dataset)
|
||||||
|
|
||||||
|
|
||||||
# ============== JSON Storage Tests (raw text + tokenizer) ==============
|
|
||||||
|
|
||||||
|
|
||||||
def _make_tokenizer_fn(tokenizer):
|
|
||||||
"""Wrap tokenizer.encode() as a str -> List[int] callable."""
|
|
||||||
return lambda text: tokenizer.encode(text, add_special_tokens=False)
|
|
||||||
|
|
||||||
|
|
||||||
def test_seq_dataset_from_json_text(base_test_env):
|
|
||||||
"""Test loading SEQ dataset from raw-text JSON with tokenizer"""
|
|
||||||
tokenizer = base_test_env["tokenizer"]
|
|
||||||
tokenizer_fn = _make_tokenizer_fn(tokenizer)
|
|
||||||
test_dir = base_test_env["test_dir"]
|
|
||||||
data_dir = os.path.join(test_dir, "json_text")
|
|
||||||
os.makedirs(data_dir, exist_ok=True)
|
|
||||||
|
|
||||||
texts = [
|
|
||||||
"hello world this is a test sentence for tokenizer",
|
|
||||||
"another sentence with different words and tokens",
|
|
||||||
"machine learning is fascinating and powerful",
|
|
||||||
]
|
|
||||||
|
|
||||||
json_path = os.path.join(data_dir, "seq_data.json")
|
|
||||||
with open(json_path, "w", encoding="utf-8") as f:
|
|
||||||
json.dump({"sequence": texts}, f, ensure_ascii=False)
|
|
||||||
|
|
||||||
dataset = DatasetFactory.load(
|
|
||||||
train_type="seq",
|
|
||||||
load_path=data_dir,
|
|
||||||
window_size=16,
|
|
||||||
tokenizer=tokenizer_fn,
|
|
||||||
)
|
|
||||||
assert dataset is not None
|
|
||||||
assert len(dataset) > 0
|
|
||||||
assert dataset.count > 0
|
|
||||||
assert "sequence" in dataset.keys
|
|
||||||
|
|
||||||
item = dataset[0]
|
|
||||||
assert "input_ids" in item
|
|
||||||
assert "target_ids" in item
|
|
||||||
assert item["input_ids"].shape[0] == 16
|
|
||||||
|
|
||||||
|
|
||||||
def test_sft_dataset_from_json_text(base_test_env):
|
|
||||||
"""Test loading SFT dataset from raw-text JSON with tokenizer"""
|
|
||||||
tokenizer = base_test_env["tokenizer"]
|
|
||||||
tokenizer_fn = _make_tokenizer_fn(tokenizer)
|
|
||||||
test_dir = base_test_env["test_dir"]
|
|
||||||
data_dir = os.path.join(test_dir, "json_sft")
|
|
||||||
os.makedirs(data_dir, exist_ok=True)
|
|
||||||
|
|
||||||
texts = [
|
|
||||||
"user asks a question about the weather",
|
|
||||||
"assistant provides a helpful response to the user",
|
|
||||||
]
|
|
||||||
|
|
||||||
json_path = os.path.join(data_dir, "sft_data.json")
|
|
||||||
with open(json_path, "w", encoding="utf-8") as f:
|
|
||||||
json.dump(
|
|
||||||
{"sequence": texts, "loss_mask": texts},
|
|
||||||
f,
|
|
||||||
ensure_ascii=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
dataset = DatasetFactory.load(
|
|
||||||
train_type="sft",
|
|
||||||
load_path=data_dir,
|
|
||||||
window_size=16,
|
|
||||||
tokenizer=tokenizer_fn,
|
|
||||||
)
|
|
||||||
assert dataset is not None
|
|
||||||
assert len(dataset) > 0
|
|
||||||
|
|
||||||
item = dataset[0]
|
|
||||||
assert "loss_mask" in item
|
|
||||||
|
|
||||||
|
|
||||||
def test_json_storage_explicit_tokenizer(base_test_env):
|
|
||||||
"""Test explicit JSON storage with tokenizer"""
|
|
||||||
tokenizer = base_test_env["tokenizer"]
|
|
||||||
tokenizer_fn = _make_tokenizer_fn(tokenizer)
|
|
||||||
test_dir = base_test_env["test_dir"]
|
|
||||||
data_dir = os.path.join(test_dir, "json_explicit")
|
|
||||||
os.makedirs(data_dir, exist_ok=True)
|
|
||||||
|
|
||||||
texts = ["abcdefghijklmnopqrstuvwxyz" * 10]
|
|
||||||
|
|
||||||
json_path = os.path.join(data_dir, "data.json")
|
|
||||||
with open(json_path, "w", encoding="utf-8") as f:
|
|
||||||
json.dump({"sequence": texts}, f, ensure_ascii=False)
|
|
||||||
|
|
||||||
token_count = len(tokenizer_fn(texts[0]))
|
|
||||||
|
|
||||||
dataset = DatasetFactory.load(
|
|
||||||
train_type="seq",
|
|
||||||
load_path=data_dir,
|
|
||||||
window_size=32,
|
|
||||||
storage_type="json",
|
|
||||||
tokenizer=tokenizer_fn,
|
|
||||||
)
|
|
||||||
assert dataset is not None
|
|
||||||
assert len(dataset) > 0
|
|
||||||
assert dataset.count == token_count
|
|
||||||
|
|
||||||
|
|
||||||
def test_dataset_count_property(base_test_env):
|
def test_dataset_count_property(base_test_env):
|
||||||
"""Test the count property returns correct raw token count"""
|
"""Test the count property returns correct raw token count"""
|
||||||
test_dir = base_test_env["test_dir"]
|
test_dir = base_test_env["test_dir"]
|
||||||
@@ -318,37 +213,29 @@ def test_unloaded_dataset_len():
|
|||||||
assert len(dataset) == 0
|
assert len(dataset) == 0
|
||||||
|
|
||||||
|
|
||||||
def test_base_segment_fetcher_empty():
|
def test_store_unloaded_len():
|
||||||
"""BaseSegmentFetcher with empty segments list"""
|
"""Unloaded Store has __len__ == 0"""
|
||||||
fetcher = BaseSegmentFetcher([])
|
store = H5Store()
|
||||||
assert len(fetcher) == 0
|
assert len(store) == 0
|
||||||
with pytest.raises(ValueError, match="out of bounds"):
|
assert store.keys == []
|
||||||
fetcher.fetch_data(0, 1)
|
|
||||||
|
|
||||||
|
|
||||||
def test_base_segment_fetcher_begin_equals_end(base_test_env):
|
def test_store_fetch_begin_equals_end(base_test_env):
|
||||||
"""fetch_data with begin == end returns empty tensor"""
|
"""Store.fetch with begin == end returns empty tensor"""
|
||||||
test_dir = base_test_env["test_dir"]
|
test_dir = base_test_env["test_dir"]
|
||||||
dummy = {"sequence": [torch.randint(0, 1000, (100,), dtype=torch.int64)]}
|
dummy = {"sequence": [torch.randint(0, 1000, (100,), dtype=torch.int64)]}
|
||||||
save_h5(test_dir, "empty_fetch", dummy)
|
save_h5(test_dir, "empty_fetch", dummy)
|
||||||
|
|
||||||
dataset = DatasetFactory.load("seq", test_dir, window_size=32)
|
dataset = DatasetFactory.load("seq", test_dir, window_size=32)
|
||||||
fetcher = dataset.storage._fetcher.multi_fetchers["sequence"]
|
result = dataset.storage.fetch(10, 10, "sequence")
|
||||||
result = fetcher.fetch_data(10, 10)
|
|
||||||
assert result.numel() == 0
|
assert result.numel() == 0
|
||||||
|
|
||||||
|
|
||||||
def test_multi_segment_fetcher_empty_dict():
|
def test_store_fetch_before_load():
|
||||||
"""MultiSegmentFetcher with empty dict has __len__ == 0"""
|
"""Store.fetch before load raises RuntimeError"""
|
||||||
fetcher = MultiSegmentFetcher({})
|
store = H5Store()
|
||||||
assert len(fetcher) == 0
|
|
||||||
|
|
||||||
|
|
||||||
def test_storage_fetch_before_load():
|
|
||||||
"""BaseStorage.fetch before load raises RuntimeError"""
|
|
||||||
storage = H5Storage()
|
|
||||||
with pytest.raises(RuntimeError, match="not loaded"):
|
with pytest.raises(RuntimeError, match="not loaded"):
|
||||||
storage.fetch(0, 10, "sequence")
|
store.fetch(0, 10, "sequence")
|
||||||
|
|
||||||
|
|
||||||
def test_detect_format_nonexistent_path():
|
def test_detect_format_nonexistent_path():
|
||||||
@@ -367,54 +254,192 @@ def test_detect_format_unsupported_file(base_test_env):
|
|||||||
detect_format(path)
|
detect_format(path)
|
||||||
|
|
||||||
|
|
||||||
def test_create_storage_invalid_type():
|
def test_create_store_invalid_type():
|
||||||
"""StorageFactory.create raises ValueError for unknown type"""
|
"""StoreFactory.create raises ValueError for unknown type"""
|
||||||
with pytest.raises(ValueError, match="Unknown component"):
|
with pytest.raises(ValueError, match="Unknown component"):
|
||||||
StorageFactory.create("parquet")
|
StoreFactory.create("parquet")
|
||||||
|
|
||||||
|
|
||||||
def test_json_pretokenized_without_tokenizer(base_test_env):
|
def test_store_multi_segment_concat(base_test_env):
|
||||||
"""Pre-tokenized JSON (List[List[int]]) loads without tokenizer"""
|
"""Multi-segment H5 data is concatenated into single tensor at load time"""
|
||||||
|
import os
|
||||||
|
|
||||||
test_dir = base_test_env["test_dir"]
|
test_dir = base_test_env["test_dir"]
|
||||||
data_dir = os.path.join(test_dir, "json_pretok")
|
data_dir = os.path.join(test_dir, "multi_seg")
|
||||||
os.makedirs(data_dir, exist_ok=True)
|
os.makedirs(data_dir, exist_ok=True)
|
||||||
|
|
||||||
json_path = os.path.join(data_dir, "data.json")
|
|
||||||
with open(json_path, "w", encoding="utf-8") as f:
|
|
||||||
json.dump({"sequence": [[1, 2, 3, 4, 5], [6, 7, 8, 9, 10]]}, f)
|
|
||||||
|
|
||||||
dataset = DatasetFactory.load("seq", data_dir, window_size=4, storage_type="json")
|
|
||||||
assert len(dataset) > 0
|
|
||||||
assert dataset.count == 10
|
|
||||||
|
|
||||||
item = dataset[0]
|
|
||||||
assert item["input_ids"].tolist() == [1, 2, 3, 4]
|
|
||||||
assert item["target_ids"].tolist() == [2, 3, 4, 5]
|
|
||||||
|
|
||||||
|
|
||||||
def test_load_json_skips_config_file(base_test_env):
|
|
||||||
"""load_json skips scalar-value config files"""
|
|
||||||
test_dir = base_test_env["test_dir"]
|
|
||||||
with open(os.path.join(test_dir, "config.json"), "w") as f:
|
|
||||||
json.dump({"vocab_size": 1000, "dim": 16}, f)
|
|
||||||
|
|
||||||
with open(os.path.join(test_dir, "data.json"), "w") as f:
|
|
||||||
json.dump({"sequence": [[1, 2, 3, 4, 5]]}, f)
|
|
||||||
|
|
||||||
result = load_json(test_dir)
|
|
||||||
assert "sequence" in result
|
|
||||||
assert "vocab_size" not in result
|
|
||||||
assert len(result["sequence"]) == 1
|
|
||||||
|
|
||||||
|
|
||||||
def test_base_segment_fetcher_multi_segment():
|
|
||||||
"""fetch_data across multiple segment boundaries"""
|
|
||||||
segs = [
|
segs = [
|
||||||
torch.tensor([1, 2, 3]),
|
torch.tensor([1, 2, 3]),
|
||||||
torch.tensor([4, 5, 6, 7]),
|
torch.tensor([4, 5, 6, 7]),
|
||||||
torch.tensor([8, 9]),
|
torch.tensor([8, 9]),
|
||||||
]
|
]
|
||||||
fetcher = BaseSegmentFetcher(segs)
|
save_h5(data_dir, "data", {"sequence": segs})
|
||||||
assert len(fetcher) == 9
|
|
||||||
result = fetcher.fetch_data(2, 7)
|
store = StoreFactory.create("h5")
|
||||||
|
store.load(data_dir)
|
||||||
|
assert len(store) == 9
|
||||||
|
result = store.fetch(2, 7, "sequence")
|
||||||
assert result.tolist() == [3, 4, 5, 6, 7]
|
assert result.tolist() == [3, 4, 5, 6, 7]
|
||||||
|
|
||||||
|
|
||||||
|
def test_save_load_bin_roundtrip(base_test_env):
|
||||||
|
"""save_bin + load_bin roundtrip preserves data"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
|
||||||
|
data = {
|
||||||
|
"sequence": [torch.tensor([1, 2, 3, 4, 5], dtype=torch.int64)],
|
||||||
|
"loss_mask": [torch.tensor([0, 1, 1, 0, 1], dtype=torch.int64)],
|
||||||
|
}
|
||||||
|
save_bin(test_dir, data)
|
||||||
|
result = load_bin(test_dir)
|
||||||
|
|
||||||
|
assert "sequence" in result
|
||||||
|
assert "loss_mask" in result
|
||||||
|
assert result["sequence"][0].tolist() == [1, 2, 3, 4, 5]
|
||||||
|
assert result["loss_mask"][0].tolist() == [0, 1, 1, 0, 1]
|
||||||
|
|
||||||
|
|
||||||
|
def test_mmap_store_load_and_fetch(base_test_env):
|
||||||
|
"""MmapStore loads bin data and fetches correctly"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
|
||||||
|
data = {
|
||||||
|
"sequence": [torch.randint(0, 1000, (200,), dtype=torch.int64)],
|
||||||
|
}
|
||||||
|
save_bin(test_dir, data)
|
||||||
|
|
||||||
|
store = StoreFactory.create("bin")
|
||||||
|
store.load(test_dir)
|
||||||
|
assert len(store) == 200
|
||||||
|
assert "sequence" in store.keys
|
||||||
|
|
||||||
|
result = store.fetch(10, 20, "sequence")
|
||||||
|
assert result.tolist() == data["sequence"][0][10:20].tolist()
|
||||||
|
|
||||||
|
|
||||||
|
def test_mmap_dataset_load(base_test_env):
|
||||||
|
"""DatasetFactory.load auto-detects bin format"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
|
||||||
|
data = {
|
||||||
|
"sequence": [torch.randint(0, 1000, (200,), dtype=torch.int64)],
|
||||||
|
}
|
||||||
|
save_bin(test_dir, data)
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load("seq", test_dir, window_size=64)
|
||||||
|
assert len(dataset) > 0
|
||||||
|
assert dataset.count == 200
|
||||||
|
assert dataset[0]["input_ids"].shape[0] == 64
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_empty_key():
|
||||||
|
"""_normalize with empty tensor list does not crash"""
|
||||||
|
store = H5Store()
|
||||||
|
store._normalize({"sequence": []})
|
||||||
|
assert len(store) == 0
|
||||||
|
assert store.keys == ["sequence"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_mixed_empty_key():
|
||||||
|
"""_normalize with empty + non-empty keys returns min=0"""
|
||||||
|
store = H5Store()
|
||||||
|
store._normalize({"sequence": [torch.tensor([1, 2, 3])], "loss_mask": []})
|
||||||
|
assert len(store) == 0
|
||||||
|
assert set(store.keys) == {"sequence", "loss_mask"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_grpo_dataset_dtype(base_test_env):
|
||||||
|
"""GRPODataset returns correct dtypes"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
|
||||||
|
seq_len = 100
|
||||||
|
data = {
|
||||||
|
"prompts": [torch.randint(0, 100, (seq_len,), dtype=torch.int32)],
|
||||||
|
"responses": [torch.randint(0, 100, (seq_len,), dtype=torch.int32)],
|
||||||
|
"masks": [torch.ones(seq_len, dtype=torch.int32)],
|
||||||
|
"rewards": [torch.ones(seq_len, dtype=torch.float32)],
|
||||||
|
}
|
||||||
|
save_h5(test_dir, "grpo_dtype", data)
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load("grpo", test_dir, window_size=32)
|
||||||
|
item = dataset[0]
|
||||||
|
|
||||||
|
assert item["prompts"].dtype == torch.long
|
||||||
|
assert item["responses"].dtype == torch.long
|
||||||
|
assert item["masks"].dtype == torch.bool
|
||||||
|
assert item["rewards"].dtype == torch.float32
|
||||||
|
|
||||||
|
|
||||||
|
def test_grpo_dataset_load(base_test_env):
|
||||||
|
"""GRPODataset loads and returns correct keys"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
seq_len = 200
|
||||||
|
data = {
|
||||||
|
"prompts": [torch.randint(0, 1000, (seq_len,), dtype=torch.int64)],
|
||||||
|
"responses": [torch.randint(0, 1000, (seq_len,), dtype=torch.int64)],
|
||||||
|
"masks": [torch.ones(seq_len, dtype=torch.int64)],
|
||||||
|
"rewards": [torch.rand(seq_len, dtype=torch.float32)],
|
||||||
|
}
|
||||||
|
save_h5(test_dir, "grpo_test", data)
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load("grpo", test_dir, window_size=64)
|
||||||
|
assert len(dataset) > 0
|
||||||
|
item = dataset[0]
|
||||||
|
assert "prompts" in item
|
||||||
|
assert "responses" in item
|
||||||
|
assert "masks" in item
|
||||||
|
assert "rewards" in item
|
||||||
|
assert item["prompts"].shape[0] == 64
|
||||||
|
assert item["responses"].shape[0] == 64
|
||||||
|
|
||||||
|
|
||||||
|
def test_detect_format_bin_dir(base_test_env):
|
||||||
|
"""detect_format returns 'bin' for directory with .bin + meta.json"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
save_bin(test_dir, {"sequence": [torch.randint(0, 100, (10,))]})
|
||||||
|
assert detect_format(test_dir) == "bin"
|
||||||
|
|
||||||
|
|
||||||
|
def test_store_fetch_multi_key(base_test_env):
|
||||||
|
"""Store.fetch with List[str] returns Dict[str, Tensor]"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
save_h5(
|
||||||
|
test_dir,
|
||||||
|
"multi_key",
|
||||||
|
{
|
||||||
|
"sequence": [torch.randint(0, 100, (100,), dtype=torch.int64)],
|
||||||
|
"loss_mask": [torch.ones(100, dtype=torch.int64)],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
store = StoreFactory.create("h5")
|
||||||
|
store.load(test_dir)
|
||||||
|
result = store.fetch(10, 20, ["sequence", "loss_mask"])
|
||||||
|
assert isinstance(result, dict)
|
||||||
|
assert result["sequence"].shape[0] == 10
|
||||||
|
assert result["loss_mask"].shape[0] == 10
|
||||||
|
|
||||||
|
|
||||||
|
def test_store_fetch_out_of_bounds(base_test_env):
|
||||||
|
"""Store.fetch raises ValueError for out-of-bounds indices"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
save_h5(test_dir, "bounds", {"sequence": [torch.randint(0, 100, (50,))]})
|
||||||
|
|
||||||
|
store = StoreFactory.create("h5")
|
||||||
|
store.load(test_dir)
|
||||||
|
with pytest.raises(ValueError, match="out of bounds"):
|
||||||
|
store.fetch(-1, 10, "sequence")
|
||||||
|
with pytest.raises(ValueError, match="out of bounds"):
|
||||||
|
store.fetch(0, 51, "sequence")
|
||||||
|
with pytest.raises(ValueError, match="out of bounds"):
|
||||||
|
store.fetch(50, 50, "sequence")
|
||||||
|
|
||||||
|
|
||||||
|
def test_dataset_load_explicit_storage_type(base_test_env):
|
||||||
|
"""DatasetFactory.load with explicit storage_type bypasses auto-detect"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
save_h5(test_dir, "explicit", {"sequence": [torch.randint(0, 100, (200,))]})
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load("seq", test_dir, window_size=64, storage_type="h5")
|
||||||
|
assert len(dataset) > 0
|
||||||
|
assert dataset.count == 200
|
||||||
|
|||||||
@@ -0,0 +1,286 @@
|
|||||||
|
"""Unit tests for protocol builders, StopChecker, GenContext, StopInfo."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||||
|
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||||
|
from astrai.inference.api.protocol import GenContext, StopChecker, StopInfo
|
||||||
|
from astrai.inference.engine import GenerationRequest
|
||||||
|
|
||||||
|
|
||||||
|
def _make_ctx(**kwargs):
|
||||||
|
defaults = {
|
||||||
|
"resp_id": "test-123",
|
||||||
|
"created": 1000,
|
||||||
|
"model": "test-model",
|
||||||
|
"prompt_tokens": 10,
|
||||||
|
"completion_tokens": 5,
|
||||||
|
}
|
||||||
|
defaults.update(kwargs)
|
||||||
|
return GenContext(**defaults)
|
||||||
|
|
||||||
|
|
||||||
|
def _sse_payloads(events):
|
||||||
|
payloads = []
|
||||||
|
for chunk in events:
|
||||||
|
for line in chunk.strip().split("\n"):
|
||||||
|
if line.startswith("data: "):
|
||||||
|
try:
|
||||||
|
payloads.append(json.loads(line[6:]))
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
pass
|
||||||
|
return payloads
|
||||||
|
|
||||||
|
|
||||||
|
class TestStopChecker:
|
||||||
|
def test_check_finds_match(self):
|
||||||
|
sc = StopChecker(["stop", "end"])
|
||||||
|
assert sc.check("hello stop world") == "stop"
|
||||||
|
|
||||||
|
def test_check_returns_none_when_no_match(self):
|
||||||
|
sc = StopChecker(["stop"])
|
||||||
|
assert sc.check("hello world") is None
|
||||||
|
|
||||||
|
def test_check_empty_sequences(self):
|
||||||
|
sc = StopChecker([])
|
||||||
|
assert sc.check("hello") is None
|
||||||
|
|
||||||
|
|
||||||
|
class TestGenContext:
|
||||||
|
def test_defaults(self):
|
||||||
|
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
|
||||||
|
assert ctx.completion_tokens == 0
|
||||||
|
|
||||||
|
def test_fields_mutable(self):
|
||||||
|
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
|
||||||
|
ctx.completion_tokens = 42
|
||||||
|
assert ctx.completion_tokens == 42
|
||||||
|
|
||||||
|
|
||||||
|
class TestStopInfo:
|
||||||
|
def test_defaults(self):
|
||||||
|
s = StopInfo()
|
||||||
|
assert s.matched is None
|
||||||
|
assert s.body == ""
|
||||||
|
assert s.yielded == ""
|
||||||
|
|
||||||
|
def test_with_values(self):
|
||||||
|
s = StopInfo(matched="stop", body="hello stop", yielded="hello ")
|
||||||
|
assert s.matched == "stop"
|
||||||
|
assert s.body == "hello stop"
|
||||||
|
assert s.yielded == "hello "
|
||||||
|
|
||||||
|
|
||||||
|
class TestOpenAIResponseBuilder:
|
||||||
|
@pytest.fixture
|
||||||
|
def builder(self):
|
||||||
|
builder = OpenAIResponseBuilder()
|
||||||
|
req = MagicMock()
|
||||||
|
req.messages = [MagicMock(role="user", content="Hello")]
|
||||||
|
req.stop = None
|
||||||
|
req.model = "astrai"
|
||||||
|
engine = MagicMock()
|
||||||
|
engine.tokenizer.apply_chat_template.return_value = "Hello"
|
||||||
|
builder.prepare(req, engine)
|
||||||
|
return builder
|
||||||
|
|
||||||
|
def test_prepare_returns_prompt_ctx_stops(self, builder):
|
||||||
|
req = MagicMock()
|
||||||
|
req.messages = [MagicMock(role="user", content="Hi")]
|
||||||
|
req.stop = ["END"]
|
||||||
|
req.model = "gpt"
|
||||||
|
engine = MagicMock()
|
||||||
|
engine.tokenizer.apply_chat_template.return_value = "Hi"
|
||||||
|
prompt, ctx, stops = builder.prepare(req, engine)
|
||||||
|
assert prompt == "Hi"
|
||||||
|
assert ctx.model == "gpt"
|
||||||
|
assert ctx.prompt_tokens == 0
|
||||||
|
assert stops == ["END"]
|
||||||
|
|
||||||
|
def test_prepare_no_stop_returns_empty_list(self, builder):
|
||||||
|
req = MagicMock()
|
||||||
|
req.messages = []
|
||||||
|
req.stop = None
|
||||||
|
req.model = "x"
|
||||||
|
engine = MagicMock()
|
||||||
|
engine.tokenizer.apply_chat_template.return_value = ""
|
||||||
|
_, _, stops = builder.prepare(req, engine)
|
||||||
|
assert stops == []
|
||||||
|
|
||||||
|
def test_format_stream_start(self, builder):
|
||||||
|
ctx = _make_ctx()
|
||||||
|
events = builder.format_stream_start(ctx)
|
||||||
|
payloads = _sse_payloads(events)
|
||||||
|
assert len(payloads) == 1
|
||||||
|
p = payloads[0]
|
||||||
|
assert p["object"] == "chat.completion.chunk"
|
||||||
|
assert p["choices"][0]["delta"]["role"] == "assistant"
|
||||||
|
assert p["choices"][0]["finish_reason"] is None
|
||||||
|
|
||||||
|
def test_format_chunk(self, builder):
|
||||||
|
event = builder.format_chunk("hello")
|
||||||
|
payload = json.loads(event.split("data: ", 1)[1])
|
||||||
|
assert payload["choices"][0]["delta"]["content"] == "hello"
|
||||||
|
assert payload["choices"][0]["finish_reason"] is None
|
||||||
|
|
||||||
|
def test_format_stream_end(self, builder):
|
||||||
|
ctx = _make_ctx(completion_tokens=5)
|
||||||
|
stop = StopInfo(matched="stop")
|
||||||
|
events = builder.format_stream_end(ctx, stop)
|
||||||
|
payloads = _sse_payloads(events)
|
||||||
|
finish = payloads[0]
|
||||||
|
assert finish["choices"][0]["finish_reason"] == "stop"
|
||||||
|
usage = payloads[1]
|
||||||
|
assert usage["completion_tokens"] == 5
|
||||||
|
assert usage["total_tokens"] == 15
|
||||||
|
|
||||||
|
def test_format_response(self, builder):
|
||||||
|
ctx = _make_ctx()
|
||||||
|
stop = StopInfo()
|
||||||
|
resp = builder.format_response(ctx, "hello", stop)
|
||||||
|
assert resp["object"] == "chat.completion"
|
||||||
|
assert resp["choices"][0]["message"]["content"] == "hello"
|
||||||
|
assert resp["usage"]["prompt_tokens"] == 10
|
||||||
|
|
||||||
|
|
||||||
|
class TestAnthropicResponseBuilder:
|
||||||
|
@pytest.fixture
|
||||||
|
def builder(self):
|
||||||
|
builder = AnthropicResponseBuilder()
|
||||||
|
req = MagicMock()
|
||||||
|
req.messages = [MagicMock(role="user", content="Hello")]
|
||||||
|
req.model = "claude"
|
||||||
|
engine = MagicMock()
|
||||||
|
engine.tokenizer.apply_chat_template.return_value = "Hello"
|
||||||
|
req.system = None
|
||||||
|
builder.prepare(req, engine)
|
||||||
|
return builder
|
||||||
|
|
||||||
|
def test_prepare_messages(self, builder):
|
||||||
|
req = MagicMock()
|
||||||
|
req.messages = [MagicMock(role="user", content="Hi")]
|
||||||
|
req.model = "claude"
|
||||||
|
req.system = None
|
||||||
|
req.stop_sequences = None
|
||||||
|
engine = MagicMock()
|
||||||
|
engine.tokenizer.apply_chat_template.return_value = "Hi"
|
||||||
|
prompt, ctx, stops = builder.prepare(req, engine)
|
||||||
|
assert prompt == "Hi"
|
||||||
|
assert stops == []
|
||||||
|
|
||||||
|
def test_prepare_with_stop_sequences(self, builder):
|
||||||
|
req = MagicMock()
|
||||||
|
req.messages = []
|
||||||
|
req.model = "x"
|
||||||
|
req.stop_sequences = ["stop", "end"]
|
||||||
|
req.system = None
|
||||||
|
engine = MagicMock()
|
||||||
|
engine.tokenizer.apply_chat_template.return_value = ""
|
||||||
|
_, _, stops = builder.prepare(req, engine)
|
||||||
|
assert stops == ["stop", "end"]
|
||||||
|
|
||||||
|
def test_format_stream_start(self, builder):
|
||||||
|
ctx = _make_ctx(prompt_tokens=3)
|
||||||
|
events = builder.format_stream_start(ctx)
|
||||||
|
payloads = _sse_payloads(events)
|
||||||
|
assert len(payloads) == 2
|
||||||
|
assert payloads[0]["type"] == "message_start"
|
||||||
|
assert payloads[0]["message"]["usage"]["input_tokens"] == 3
|
||||||
|
assert payloads[1]["type"] == "content_block_start"
|
||||||
|
|
||||||
|
def test_format_chunk(self, builder):
|
||||||
|
event = builder.format_chunk("tok")
|
||||||
|
payload = json.loads(event.split("data: ", 1)[1])
|
||||||
|
assert payload["type"] == "content_block_delta"
|
||||||
|
assert payload["delta"]["text"] == "tok"
|
||||||
|
|
||||||
|
def test_format_stream_end_no_stop(self, builder):
|
||||||
|
ctx = _make_ctx(completion_tokens=3)
|
||||||
|
stop = StopInfo()
|
||||||
|
events = builder.format_stream_end(ctx, stop)
|
||||||
|
payloads = _sse_payloads(events)
|
||||||
|
# content_block_stop, message_delta, message_stop
|
||||||
|
types = [p["type"] for p in payloads]
|
||||||
|
assert types == ["content_block_stop", "message_delta", "message_stop"]
|
||||||
|
assert payloads[1]["delta"]["stop_reason"] == "end_turn"
|
||||||
|
|
||||||
|
def test_format_stream_end_with_stop_trims_and_emits_remaining(self, builder):
|
||||||
|
ctx = _make_ctx(completion_tokens=7)
|
||||||
|
stop = StopInfo(
|
||||||
|
matched="END",
|
||||||
|
body="Hello world END extra",
|
||||||
|
yielded="Hello ",
|
||||||
|
)
|
||||||
|
events = builder.format_stream_end(ctx, stop)
|
||||||
|
payloads = _sse_payloads(events)
|
||||||
|
# unyielded delta, content_block_stop, message_delta, message_stop
|
||||||
|
types = [p["type"] for p in payloads]
|
||||||
|
assert types == [
|
||||||
|
"content_block_delta",
|
||||||
|
"content_block_stop",
|
||||||
|
"message_delta",
|
||||||
|
"message_stop",
|
||||||
|
]
|
||||||
|
assert payloads[0]["delta"]["text"] == "world "
|
||||||
|
assert payloads[2]["delta"]["stop_reason"] == "stop_sequence"
|
||||||
|
assert payloads[2]["delta"]["stop_sequence"] == "END"
|
||||||
|
|
||||||
|
def test_format_stream_end_stop_trimmed_already_yielded(self, builder):
|
||||||
|
ctx = _make_ctx()
|
||||||
|
stop = StopInfo(
|
||||||
|
matched="END",
|
||||||
|
body="Hello END",
|
||||||
|
yielded="Hello ",
|
||||||
|
)
|
||||||
|
events = builder.format_stream_end(ctx, stop)
|
||||||
|
payloads = _sse_payloads(events)
|
||||||
|
# No unyielded delta (everything already sent)
|
||||||
|
types = [p["type"] for p in payloads]
|
||||||
|
assert types == ["content_block_stop", "message_delta", "message_stop"]
|
||||||
|
|
||||||
|
def test_format_response_with_stop_trims_content(self, builder):
|
||||||
|
ctx = _make_ctx()
|
||||||
|
stop = StopInfo(matched="STOP", body="text STOP extra", yielded="text ")
|
||||||
|
resp = builder.format_response(ctx, "text STOP extra", stop)
|
||||||
|
assert resp["content"][0]["text"] == "text "
|
||||||
|
assert resp["stop_reason"] == "stop_sequence"
|
||||||
|
assert resp["stop_sequence"] == "STOP"
|
||||||
|
|
||||||
|
def test_format_response_no_stop(self, builder):
|
||||||
|
ctx = _make_ctx()
|
||||||
|
stop = StopInfo()
|
||||||
|
resp = builder.format_response(ctx, "full text", stop)
|
||||||
|
assert resp["content"][0]["text"] == "full text"
|
||||||
|
assert resp["stop_reason"] == "end_turn"
|
||||||
|
|
||||||
|
|
||||||
|
class TestGenerationRequestValidation:
|
||||||
|
def test_valid_params(self):
|
||||||
|
gr = GenerationRequest(
|
||||||
|
messages=[{"role": "user", "content": "hi"}],
|
||||||
|
top_k=50,
|
||||||
|
top_p=0.9,
|
||||||
|
temperature=0.7,
|
||||||
|
)
|
||||||
|
assert gr.top_k == 50
|
||||||
|
|
||||||
|
def test_invalid_top_p_raises(self):
|
||||||
|
with pytest.raises(ValueError, match="top_p"):
|
||||||
|
GenerationRequest(messages=[{"role": "user", "content": "hi"}], top_p=1.5)
|
||||||
|
|
||||||
|
def test_invalid_top_k_raises(self):
|
||||||
|
with pytest.raises(ValueError, match="top_k"):
|
||||||
|
GenerationRequest(messages=[{"role": "user", "content": "hi"}], top_k=-1)
|
||||||
|
|
||||||
|
def test_invalid_temperature_raises(self):
|
||||||
|
with pytest.raises(ValueError, match="temperature"):
|
||||||
|
GenerationRequest(
|
||||||
|
messages=[{"role": "user", "content": "hi"}], temperature=-0.1
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_top_k_zero_valid(self):
|
||||||
|
gr = GenerationRequest(messages=[{"role": "user", "content": "hi"}], top_k=0)
|
||||||
|
assert gr.top_k == 0
|
||||||
@@ -173,3 +173,21 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
|
|||||||
for stats in results["stats"]:
|
for stats in results["stats"]:
|
||||||
assert "total_tasks" in stats
|
assert "total_tasks" in stats
|
||||||
assert stats["total_tasks"] >= 0
|
assert stats["total_tasks"] >= 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefill_skips_fully_cached_tasks(mock_model_and_tokenizer):
|
||||||
|
"""Tasks whose entire prompt is cached skip the prefill phase."""
|
||||||
|
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
||||||
|
|
||||||
|
with patch("astrai.inference.core.scheduler.AutoModel"):
|
||||||
|
with patch("astrai.inference.core.scheduler.AutoTokenizer"):
|
||||||
|
scheduler = InferenceScheduler(
|
||||||
|
model=mock_model,
|
||||||
|
tokenizer=mock_tokenizer,
|
||||||
|
max_batch_size=4,
|
||||||
|
device="cpu",
|
||||||
|
)
|
||||||
|
|
||||||
|
task_id = scheduler.add_task("short prompt", stream_callback=lambda t: None)
|
||||||
|
scheduler.stop()
|
||||||
|
assert task_id.startswith("task_")
|
||||||
|
|||||||
@@ -0,0 +1,355 @@
|
|||||||
|
import tempfile
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
from astrai.model import AutoRegressiveLM
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
from astrai.model.components.lora import (
|
||||||
|
LoRAConfig,
|
||||||
|
LoRALinear,
|
||||||
|
_collect_lora_info,
|
||||||
|
_get_lora_count,
|
||||||
|
inject_lora,
|
||||||
|
load_lora,
|
||||||
|
merge_lora,
|
||||||
|
save_lora,
|
||||||
|
)
|
||||||
|
|
||||||
|
MODEL_KWARGS = dict(
|
||||||
|
vocab_size=1000,
|
||||||
|
dim=64,
|
||||||
|
n_heads=4,
|
||||||
|
n_kv_heads=2,
|
||||||
|
dim_ffn=128,
|
||||||
|
n_layers=2,
|
||||||
|
max_len=32,
|
||||||
|
norm_eps=1e-5,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_model(**kwargs):
|
||||||
|
kw = {**MODEL_KWARGS, **kwargs}
|
||||||
|
config = AutoRegressiveLMConfig(**kw)
|
||||||
|
model = AutoRegressiveLM(config)
|
||||||
|
model.eval()
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def test_loralinear_init():
|
||||||
|
base = Linear(64, 128)
|
||||||
|
lora = LoRALinear(base, r=8, alpha=16)
|
||||||
|
|
||||||
|
assert lora.weight is base.weight
|
||||||
|
assert not lora.weight.requires_grad
|
||||||
|
assert lora.lora_A.shape == (8, 64)
|
||||||
|
assert lora.lora_B.shape == (128, 8)
|
||||||
|
assert lora.scaling == 2.0
|
||||||
|
assert not lora._merged
|
||||||
|
assert lora.lora_A.requires_grad
|
||||||
|
assert lora.lora_B.requires_grad
|
||||||
|
|
||||||
|
|
||||||
|
def test_loralinear_forward_init_zero_delta():
|
||||||
|
base = Linear(4, 4)
|
||||||
|
with torch.no_grad():
|
||||||
|
base.weight.zero_()
|
||||||
|
|
||||||
|
x = torch.randn(2, 4)
|
||||||
|
lora = LoRALinear(base, r=2, alpha=2)
|
||||||
|
base_out = base(x)
|
||||||
|
lora_out = lora(x)
|
||||||
|
|
||||||
|
assert torch.allclose(base_out, lora_out)
|
||||||
|
|
||||||
|
|
||||||
|
def test_loralinear_forward_with_delta():
|
||||||
|
base = Linear(4, 4)
|
||||||
|
with torch.no_grad():
|
||||||
|
base.weight.zero_()
|
||||||
|
|
||||||
|
x = torch.randn(2, 4)
|
||||||
|
lora = LoRALinear(base, r=2, alpha=2)
|
||||||
|
base_out = base(x)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
lora.lora_B.fill_(1.0)
|
||||||
|
|
||||||
|
lora_out = lora(x)
|
||||||
|
assert not torch.allclose(base_out, lora_out)
|
||||||
|
|
||||||
|
|
||||||
|
def test_loralinear_merge():
|
||||||
|
base = Linear(4, 4)
|
||||||
|
with torch.no_grad():
|
||||||
|
base.weight.zero_()
|
||||||
|
|
||||||
|
x = torch.randn(2, 4)
|
||||||
|
lora = LoRALinear(base, r=2, alpha=2)
|
||||||
|
with torch.no_grad():
|
||||||
|
lora.lora_B.fill_(1.0)
|
||||||
|
|
||||||
|
out_before = lora(x).clone()
|
||||||
|
lora.merge()
|
||||||
|
out_after = lora(x)
|
||||||
|
|
||||||
|
torch.testing.assert_close(out_before, out_after)
|
||||||
|
assert lora._merged
|
||||||
|
assert not hasattr(lora, "lora_A")
|
||||||
|
|
||||||
|
|
||||||
|
def test_loralinear_merge_is_idempotent():
|
||||||
|
base = Linear(4, 4)
|
||||||
|
with torch.no_grad():
|
||||||
|
base.weight.zero_()
|
||||||
|
|
||||||
|
lora = LoRALinear(base, r=2, alpha=2)
|
||||||
|
with torch.no_grad():
|
||||||
|
lora.lora_B.fill_(1.0)
|
||||||
|
|
||||||
|
lora.merge()
|
||||||
|
lora.merge()
|
||||||
|
|
||||||
|
|
||||||
|
def test_inject_lora_default_target():
|
||||||
|
model = _make_model()
|
||||||
|
n_before = sum(1 for m in model.modules() if isinstance(m, Linear))
|
||||||
|
|
||||||
|
inject_lora(model, r=4, alpha=8)
|
||||||
|
|
||||||
|
lora_count = _get_lora_count(model)
|
||||||
|
assert lora_count > 0
|
||||||
|
assert lora_count < n_before
|
||||||
|
|
||||||
|
|
||||||
|
def test_inject_lora_ffn():
|
||||||
|
model = _make_model()
|
||||||
|
from astrai.model.components.lora import TARGET_MODULES_FFN
|
||||||
|
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules=TARGET_MODULES_FFN)
|
||||||
|
assert _get_lora_count(model) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_inject_lora_returns_config():
|
||||||
|
model = _make_model()
|
||||||
|
cfg = inject_lora(model, r=8, alpha=32)
|
||||||
|
assert isinstance(cfg, LoRAConfig)
|
||||||
|
assert cfg.r == 8
|
||||||
|
assert cfg.alpha == 32
|
||||||
|
|
||||||
|
|
||||||
|
def test_inject_lora_no_matching_targets_warns(caplog):
|
||||||
|
model = _make_model()
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"nonexistent"})
|
||||||
|
assert "No LoRA layers injected" in caplog.text
|
||||||
|
|
||||||
|
|
||||||
|
def test_inject_lora_preserves_base_output():
|
||||||
|
model = _make_model()
|
||||||
|
x = torch.randint(0, 1000, (2, 16))
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
out_before = model(x)["logits"].clone()
|
||||||
|
|
||||||
|
inject_lora(model, r=4, alpha=8)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
out_after = model(x)["logits"]
|
||||||
|
|
||||||
|
torch.testing.assert_close(out_before, out_after)
|
||||||
|
|
||||||
|
|
||||||
|
def test_inject_lora_does_not_reinject():
|
||||||
|
model = _make_model()
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
first_count = _get_lora_count(model)
|
||||||
|
|
||||||
|
inject_lora(model, r=2, alpha=4, target_modules={"q_proj"})
|
||||||
|
assert _get_lora_count(model) == first_count
|
||||||
|
|
||||||
|
|
||||||
|
def test_inject_lora_adds_new_modules():
|
||||||
|
model = _make_model()
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
first = _get_lora_count(model)
|
||||||
|
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"v_proj"})
|
||||||
|
assert _get_lora_count(model) > first
|
||||||
|
|
||||||
|
|
||||||
|
def test_inject_lora_on_mla_model():
|
||||||
|
model = _make_model(
|
||||||
|
attn_type="mla", kv_lora_rank=16, qk_nope_head_dim=16, qk_rope_head_dim=16
|
||||||
|
)
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"q_proj", "o_proj"})
|
||||||
|
assert _get_lora_count(model) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_inject_lora_on_moe_model():
|
||||||
|
model = _make_model(
|
||||||
|
ffn_type="moe",
|
||||||
|
n_routed_experts=4,
|
||||||
|
n_shared_experts=1,
|
||||||
|
n_activated_experts=2,
|
||||||
|
dim_ffn=32,
|
||||||
|
)
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"up", "gate", "down"})
|
||||||
|
assert _get_lora_count(model) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_state_dict_key_format():
|
||||||
|
model = _make_model()
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
|
||||||
|
sd = model.state_dict()
|
||||||
|
assert "layers.0.attention.q_proj.weight" in sd
|
||||||
|
assert "layers.0.attention.q_proj.lora_A" in sd
|
||||||
|
assert "layers.0.attention.q_proj.lora_B" in sd
|
||||||
|
|
||||||
|
|
||||||
|
def test_only_lora_params_trainable():
|
||||||
|
model = _make_model()
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"q_proj", "v_proj"})
|
||||||
|
|
||||||
|
for name, param in model.named_parameters():
|
||||||
|
if isinstance(name.split(".")[-1], str) and "lora" in name:
|
||||||
|
assert param.requires_grad, f"lora param should be trainable: {name}"
|
||||||
|
elif any(name.endswith(f".{t}.weight") for t in ("q_proj", "v_proj")):
|
||||||
|
assert not param.requires_grad, f"injected weight should be frozen: {name}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_state_dict_after_inject_consistent_with_original():
|
||||||
|
model = _make_model()
|
||||||
|
sd_before = {k: v for k, v in model.state_dict().items()}
|
||||||
|
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
sd_after = model.state_dict()
|
||||||
|
|
||||||
|
# original keys unchanged
|
||||||
|
for k in sd_before:
|
||||||
|
assert k in sd_after
|
||||||
|
assert sd_before[k].shape == sd_after[k].shape
|
||||||
|
|
||||||
|
# new lora keys present
|
||||||
|
lora_keys = [k for k in sd_after if "lora" in k]
|
||||||
|
assert len(lora_keys) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_save_load_roundtrip():
|
||||||
|
model = _make_model()
|
||||||
|
cfg = inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for m in model.modules():
|
||||||
|
if isinstance(m, LoRALinear):
|
||||||
|
m.lora_B.fill_(0.5)
|
||||||
|
|
||||||
|
x = torch.randint(0, 1000, (2, 16))
|
||||||
|
with torch.no_grad():
|
||||||
|
out_src = model(x)["logits"].clone()
|
||||||
|
|
||||||
|
tmpdir = tempfile.mkdtemp()
|
||||||
|
save_lora(model, tmpdir, cfg)
|
||||||
|
|
||||||
|
model2 = _make_model()
|
||||||
|
model2.load_state_dict(model.state_dict(), strict=False)
|
||||||
|
load_lora(model2, tmpdir)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
out_dst = model2(x)["logits"]
|
||||||
|
|
||||||
|
torch.testing.assert_close(out_src, out_dst)
|
||||||
|
|
||||||
|
|
||||||
|
def test_save_after_merge_raises():
|
||||||
|
model = _make_model()
|
||||||
|
cfg = inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for m in model.modules():
|
||||||
|
if isinstance(m, LoRALinear):
|
||||||
|
m.lora_B.fill_(0.5)
|
||||||
|
|
||||||
|
tmpdir = tempfile.mkdtemp()
|
||||||
|
save_lora(model, tmpdir, cfg)
|
||||||
|
merge_lora(model)
|
||||||
|
|
||||||
|
tmpdir2 = tempfile.mkdtemp()
|
||||||
|
with pytest.raises(RuntimeError, match="No LoRA parameters"):
|
||||||
|
save_lora(model, tmpdir2, cfg)
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_lora_on_already_injected():
|
||||||
|
model = _make_model()
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for m in model.modules():
|
||||||
|
if isinstance(m, LoRALinear):
|
||||||
|
m.lora_B.fill_(0.5)
|
||||||
|
|
||||||
|
tmpdir = tempfile.mkdtemp()
|
||||||
|
save_lora(model, tmpdir, LoRAConfig(r=4, alpha=8, target_modules=("q_proj",)))
|
||||||
|
|
||||||
|
model2 = _make_model()
|
||||||
|
model2.load_state_dict(model.state_dict(), strict=False)
|
||||||
|
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
|
||||||
|
# load onto already-injected model
|
||||||
|
load_lora(model2, tmpdir)
|
||||||
|
assert _get_lora_count(model2) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_lora_mismatched_r_raises():
|
||||||
|
model = _make_model()
|
||||||
|
cfg = inject_lora(model, r=8, alpha=16, target_modules={"q_proj"})
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for m in model.modules():
|
||||||
|
if isinstance(m, LoRALinear):
|
||||||
|
m.lora_B.fill_(0.5)
|
||||||
|
|
||||||
|
tmpdir = tempfile.mkdtemp()
|
||||||
|
save_lora(model, tmpdir, cfg)
|
||||||
|
|
||||||
|
model2 = _make_model()
|
||||||
|
model2.load_state_dict(model.state_dict(), strict=False)
|
||||||
|
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="size mismatch"):
|
||||||
|
load_lora(model2, tmpdir) # strict=False, only lora keys
|
||||||
|
|
||||||
|
|
||||||
|
def test_merge_preserves_output():
|
||||||
|
model = _make_model()
|
||||||
|
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for m in model.modules():
|
||||||
|
if isinstance(m, LoRALinear):
|
||||||
|
m.lora_B.fill_(0.5)
|
||||||
|
|
||||||
|
x = torch.randint(0, 1000, (2, 16))
|
||||||
|
with torch.no_grad():
|
||||||
|
out_before = model(x)["logits"].clone()
|
||||||
|
|
||||||
|
merge_lora(model)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
out_after = model(x)["logits"]
|
||||||
|
torch.testing.assert_close(out_before, out_after)
|
||||||
|
|
||||||
|
|
||||||
|
def test_merge_no_lora_warns(caplog):
|
||||||
|
model = _make_model()
|
||||||
|
merge_lora(model)
|
||||||
|
assert "No LoRA layers to merge" in caplog.text
|
||||||
|
|
||||||
|
|
||||||
|
def test_collect_lora_info():
|
||||||
|
model = _make_model()
|
||||||
|
info = _collect_lora_info(model)
|
||||||
|
assert "q_proj" in info
|
||||||
|
assert "o_proj" in info
|
||||||
|
assert "q_proj" in info # each layer has one
|
||||||
@@ -27,7 +27,7 @@ class TrainerDataset(Dataset):
|
|||||||
|
|
||||||
|
|
||||||
def create_train_config(
|
def create_train_config(
|
||||||
model: torch.nn.Module,
|
model_fn,
|
||||||
dataset: Dataset,
|
dataset: Dataset,
|
||||||
test_dir: str,
|
test_dir: str,
|
||||||
device: str,
|
device: str,
|
||||||
@@ -43,7 +43,7 @@ def create_train_config(
|
|||||||
"""Factory function to create common TrainConfig for tests.
|
"""Factory function to create common TrainConfig for tests.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model: The model to train
|
model_fn: Model factory (callable returning nn.Module)
|
||||||
dataset: Training dataset
|
dataset: Training dataset
|
||||||
test_dir: Checkpoint directory
|
test_dir: Checkpoint directory
|
||||||
device: Device type ("cuda" or "cpu")
|
device: Device type ("cuda" or "cpu")
|
||||||
@@ -70,7 +70,7 @@ def create_train_config(
|
|||||||
|
|
||||||
return TrainConfig(
|
return TrainConfig(
|
||||||
strategy=strategy,
|
strategy=strategy,
|
||||||
model=model,
|
model_fn=model_fn,
|
||||||
dataset=dataset,
|
dataset=dataset,
|
||||||
optimizer_fn=optimizer_fn,
|
optimizer_fn=optimizer_fn,
|
||||||
scheduler_fn=scheduler_fn,
|
scheduler_fn=scheduler_fn,
|
||||||
|
|||||||
@@ -106,7 +106,7 @@ def test_gradient_checkpointing_trainer_integration(base_test_env, random_datase
|
|||||||
)
|
)
|
||||||
|
|
||||||
train_config = TrainConfig(
|
train_config = TrainConfig(
|
||||||
model=base_test_env["model"],
|
model_fn=lambda: base_test_env["model"],
|
||||||
strategy="seq",
|
strategy="seq",
|
||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
optimizer_fn=optimizer_fn,
|
optimizer_fn=optimizer_fn,
|
||||||
@@ -140,7 +140,7 @@ def test_callback_integration(base_test_env, random_dataset):
|
|||||||
)
|
)
|
||||||
|
|
||||||
train_config = TrainConfig(
|
train_config = TrainConfig(
|
||||||
model=base_test_env["model"],
|
model_fn=lambda: base_test_env["model"],
|
||||||
strategy="seq",
|
strategy="seq",
|
||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
optimizer_fn=optimizer_fn,
|
optimizer_fn=optimizer_fn,
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import numpy as np
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.config.train_config import TrainConfig
|
from astrai.config.train_config import TrainConfig
|
||||||
from astrai.serialization import Checkpoint
|
|
||||||
from astrai.trainer.schedule import SchedulerFactory
|
from astrai.trainer.schedule import SchedulerFactory
|
||||||
from astrai.trainer.trainer import Trainer
|
from astrai.trainer.trainer import Trainer
|
||||||
|
|
||||||
@@ -24,7 +23,7 @@ def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
|
|||||||
strategy="seq",
|
strategy="seq",
|
||||||
optimizer_fn=optimizer_fn,
|
optimizer_fn=optimizer_fn,
|
||||||
scheduler_fn=scheduler_fn,
|
scheduler_fn=scheduler_fn,
|
||||||
model=base_test_env["model"],
|
model_fn=lambda: base_test_env["model"],
|
||||||
dataset=early_stopping_dataset,
|
dataset=early_stopping_dataset,
|
||||||
ckpt_dir=base_test_env["test_dir"],
|
ckpt_dir=base_test_env["test_dir"],
|
||||||
log_dir=os.path.join(base_test_env["test_dir"], "logs"),
|
log_dir=os.path.join(base_test_env["test_dir"], "logs"),
|
||||||
@@ -39,17 +38,20 @@ def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
|
|||||||
trainer = Trainer(train_config)
|
trainer = Trainer(train_config)
|
||||||
|
|
||||||
# Should handle early stopping gracefully
|
# Should handle early stopping gracefully
|
||||||
checkpoint = None
|
|
||||||
try:
|
try:
|
||||||
checkpoint = trainer.train()
|
trainer.train()
|
||||||
except Exception:
|
except Exception:
|
||||||
# Handle any exceptions
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
# Resume from latest checkpoint
|
||||||
load_dir = os.path.join(base_test_env["test_dir"], "epoch_0_iter_2")
|
load_dir = os.path.join(base_test_env["test_dir"], "epoch_0_iter_2")
|
||||||
checkpoint = Checkpoint.load(load_dir)
|
trainer = Trainer(train_config)
|
||||||
trainer.train(checkpoint)
|
trainer.train(resume_dir=load_dir)
|
||||||
|
|
||||||
|
# Verify checkpoint was saved at expected iteration
|
||||||
load_dir = os.path.join(base_test_env["test_dir"], "epoch_1_iter_10")
|
load_dir = os.path.join(base_test_env["test_dir"], "epoch_1_iter_10")
|
||||||
checkpoint = Checkpoint.load(load_dir)
|
import json
|
||||||
assert checkpoint.iteration == 10
|
|
||||||
|
with open(os.path.join(load_dir, "meta.json")) as f:
|
||||||
|
meta = json.load(f)
|
||||||
|
assert meta["iteration"] == 10
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ def test_different_batch_sizes(base_test_env, random_dataset, train_config_facto
|
|||||||
|
|
||||||
for batch_per_device in batch_sizes:
|
for batch_per_device in batch_sizes:
|
||||||
train_config = train_config_factory(
|
train_config = train_config_factory(
|
||||||
model=base_test_env["model"],
|
model_fn=lambda: base_test_env["model"],
|
||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
test_dir=base_test_env["test_dir"],
|
test_dir=base_test_env["test_dir"],
|
||||||
device=base_test_env["device"],
|
device=base_test_env["device"],
|
||||||
@@ -25,7 +25,7 @@ def test_gradient_accumulation(base_test_env, random_dataset, train_config_facto
|
|||||||
|
|
||||||
for grad_accum_steps in grad_accum_steps_list:
|
for grad_accum_steps in grad_accum_steps_list:
|
||||||
train_config = train_config_factory(
|
train_config = train_config_factory(
|
||||||
model=base_test_env["model"],
|
model_fn=lambda: base_test_env["model"],
|
||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
test_dir=base_test_env["test_dir"],
|
test_dir=base_test_env["test_dir"],
|
||||||
device=base_test_env["device"],
|
device=base_test_env["device"],
|
||||||
@@ -50,7 +50,7 @@ def test_memory_efficient_training(base_test_env, random_dataset, train_config_f
|
|||||||
|
|
||||||
for config in small_batch_configs:
|
for config in small_batch_configs:
|
||||||
train_config = train_config_factory(
|
train_config = train_config_factory(
|
||||||
model=base_test_env["model"],
|
model_fn=lambda: base_test_env["model"],
|
||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
test_dir=base_test_env["test_dir"],
|
test_dir=base_test_env["test_dir"],
|
||||||
device=base_test_env["device"],
|
device=base_test_env["device"],
|
||||||
|
|||||||
Reference in New Issue
Block a user