Compare commits
20
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
663ef900fc | ||
|
|
7d478a54db | ||
|
|
f3eaaef842 | ||
|
|
d655b65027 | ||
|
|
31c22dc043 | ||
|
|
17127f8b3c | ||
|
|
d7695b40e3 | ||
|
|
fc62890e70 | ||
|
|
f433672140 | ||
|
|
7e1e5b6e6a | ||
|
|
553a42702d | ||
|
|
b133fc9c07 | ||
|
|
b33250dc28 | ||
|
|
a74e5b91a3 | ||
|
|
28886e4241 | ||
|
|
9d3ccfdffc | ||
|
|
a24a7b4da5 | ||
|
|
f7df02f9a3 | ||
|
|
ee450686f3 | ||
|
|
2565755e45 |
@@ -20,7 +20,7 @@
|
|||||||
<a href="assets/docs/README-zh-CN.md">中文</a> •
|
<a href="assets/docs/README-zh-CN.md">中文</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
||||||
<a href="https://huggingface.co/ViperEk/">HuggingFace</a>
|
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<br>
|
<br>
|
||||||
@@ -241,7 +241,7 @@ For major changes, please open an issue first to discuss what you would like to
|
|||||||
|
|
||||||
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
||||||
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
||||||
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEk)
|
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEkura)
|
||||||
|
|
||||||
### License
|
### License
|
||||||
|
|
||||||
|
|||||||
@@ -27,7 +27,7 @@
|
|||||||
<a href="#chinese">中文</a> •
|
<a href="#chinese">中文</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
||||||
<a href="https://huggingface.co/ViperEk">HuggingFace</a>
|
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||||
</div>
|
</div>
|
||||||
<br>
|
<br>
|
||||||
|
|
||||||
@@ -247,7 +247,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)
|
|||||||
|
|
||||||
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
|
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
|
||||||
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
|
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
|
||||||
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEk)
|
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEkura)
|
||||||
|
|
||||||
### 许可证
|
### 许可证
|
||||||
|
|
||||||
|
|||||||
+64
-14
@@ -117,7 +117,7 @@ classDiagram
|
|||||||
+int n_epoch
|
+int n_epoch
|
||||||
+int batch_per_device
|
+int batch_per_device
|
||||||
+int grad_accum_steps
|
+int grad_accum_steps
|
||||||
+float max_grad_norm
|
+Optional[float] max_grad_norm
|
||||||
+list gradient_checkpointing_modules
|
+list gradient_checkpointing_modules
|
||||||
+int start_epoch
|
+int start_epoch
|
||||||
+int start_samples
|
+int start_samples
|
||||||
@@ -166,6 +166,13 @@ classDiagram
|
|||||||
+__getitem__(index) Dict
|
+__getitem__(index) Dict
|
||||||
}
|
}
|
||||||
|
|
||||||
|
class RecordDataset {
|
||||||
|
+Optional[Callable] processor
|
||||||
|
+load(load_path, storage_type)
|
||||||
|
+__getitem__(index)
|
||||||
|
+__len__()
|
||||||
|
}
|
||||||
|
|
||||||
class DPODataset {
|
class DPODataset {
|
||||||
+__getitem__(index) Dict
|
+__getitem__(index) Dict
|
||||||
}
|
}
|
||||||
@@ -177,13 +184,26 @@ classDiagram
|
|||||||
class Store {
|
class Store {
|
||||||
+Dict[str, List[Tensor]] _data
|
+Dict[str, List[Tensor]] _data
|
||||||
+Dict[str, List[int]] _cum
|
+Dict[str, List[int]] _cum
|
||||||
|
+Dict[str, List[int]] _offsets
|
||||||
+int _length
|
+int _length
|
||||||
|
+int _num_records
|
||||||
+keys (property)
|
+keys (property)
|
||||||
+load(path)
|
+load(path)
|
||||||
+fetch(begin, end, keys)
|
|
||||||
+__len__()
|
+__len__()
|
||||||
-_fetch_key(key, begin, end) Tensor
|
-_normalize(raw, offsets)
|
||||||
-_normalize(raw)
|
}
|
||||||
|
|
||||||
|
class Streamable {
|
||||||
|
<<mixin>>
|
||||||
|
+fetch(begin, end, keys)
|
||||||
|
-_fetch_stream_key(key, begin, end) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class Recordable {
|
||||||
|
<<mixin>>
|
||||||
|
+num_records (property)
|
||||||
|
+fetch_record(index, keys)
|
||||||
|
-_fetch_record_key(key, index) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
class H5Store {
|
class H5Store {
|
||||||
@@ -195,6 +215,13 @@ classDiagram
|
|||||||
+load(path)
|
+load(path)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
class JsonlStore {
|
||||||
|
+JsonlSource _source
|
||||||
|
+Callable _processor
|
||||||
|
+load(path, transform, processor)
|
||||||
|
+fetch_record(index, keys)
|
||||||
|
}
|
||||||
|
|
||||||
class ResumableDistributedSampler {
|
class ResumableDistributedSampler {
|
||||||
+int epoch
|
+int epoch
|
||||||
+int iter
|
+int iter
|
||||||
@@ -210,7 +237,7 @@ classDiagram
|
|||||||
+Dict _entries
|
+Dict _entries
|
||||||
+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) BaseDataset
|
+load(train_type, load_path, window_size, stride, storage_type, tokenizer_path, max_len, store) BaseDataset
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -378,6 +405,7 @@ classDiagram
|
|||||||
+List[str] paths
|
+List[str] paths
|
||||||
+str output_dir
|
+str output_dir
|
||||||
+str tokenizer_path
|
+str tokenizer_path
|
||||||
|
+AutoTokenizer tokenizer
|
||||||
+BaseMaskBuilder mask_builder
|
+BaseMaskBuilder mask_builder
|
||||||
+PackingStrategy _packer
|
+PackingStrategy _packer
|
||||||
+PositionIdStrategy _position_id
|
+PositionIdStrategy _position_id
|
||||||
@@ -385,6 +413,18 @@ classDiagram
|
|||||||
+transform(item) Optional[dict]
|
+transform(item) Optional[dict]
|
||||||
+run()
|
+run()
|
||||||
+_flush(domains, shard_idx)
|
+_flush(domains, shard_idx)
|
||||||
|
+_inject_doc_reset_position_ids(keys, mode, seqs) Dict
|
||||||
|
+_inject_continuous_position_ids(tensors, mode, seqs) Dict
|
||||||
|
+_to_tensors(keys) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class TokenizeTransform {
|
||||||
|
+PipelineConfig config
|
||||||
|
+AutoTokenizer tokenizer
|
||||||
|
+BaseMaskBuilder mask_builder
|
||||||
|
+PositionIdStrategy position_strategy
|
||||||
|
+from_config_file(path) TokenizeTransform
|
||||||
|
+apply(records) Dict[str, list]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -495,13 +535,13 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class GRPOStrategy {
|
class GRPOStrategy {
|
||||||
|
+nn.Module old_model
|
||||||
+nn.Module ref_model
|
+nn.Module ref_model
|
||||||
+float clip_eps
|
+float clip_eps
|
||||||
+float kl_coef
|
+float kl_coef
|
||||||
+int group_size
|
+int group_size
|
||||||
+int sync_interval
|
|
||||||
+compute_loss(batch) Tensor
|
+compute_loss(batch) Tensor
|
||||||
+sync_ref_model()
|
+sync_old_model()
|
||||||
}
|
}
|
||||||
|
|
||||||
class BaseScheduler {
|
class BaseScheduler {
|
||||||
@@ -551,7 +591,7 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class GradientClippingCallback {
|
class GradientClippingCallback {
|
||||||
+float max_grad_norm
|
+Optional[float] max_grad_norm
|
||||||
+on_optimizer_step(context)
|
+on_optimizer_step(context)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1064,11 +1104,18 @@ classDiagram
|
|||||||
TrainCallback <|-- MetricCallback
|
TrainCallback <|-- MetricCallback
|
||||||
BaseDataset <|-- SEQDataset
|
BaseDataset <|-- SEQDataset
|
||||||
BaseDataset <|-- SFTDataset
|
BaseDataset <|-- SFTDataset
|
||||||
BaseDataset <|-- DPODataset
|
BaseDataset <|-- RecordDataset
|
||||||
BaseDataset <|-- GRPODataset
|
RecordDataset <|-- DPODataset
|
||||||
|
RecordDataset <|-- GRPODataset
|
||||||
Store <|-- H5Store
|
Store <|-- H5Store
|
||||||
Store <|-- MmapStore
|
Store <|-- MmapStore
|
||||||
Store <|-- JsonlStore
|
Store <|-- JsonlStore
|
||||||
|
H5Store --|> Streamable
|
||||||
|
H5Store --|> Recordable
|
||||||
|
MmapStore --|> Streamable
|
||||||
|
MmapStore --|> Recordable
|
||||||
|
JsonlStore --|> Streamable
|
||||||
|
JsonlStore --|> Recordable
|
||||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||||
BaseSamplingStrategy <|-- TopKStrategy
|
BaseSamplingStrategy <|-- TopKStrategy
|
||||||
BaseSamplingStrategy <|-- TopPStrategy
|
BaseSamplingStrategy <|-- TopPStrategy
|
||||||
@@ -1143,6 +1190,9 @@ classDiagram
|
|||||||
BaseDataset o-- Store
|
BaseDataset o-- Store
|
||||||
Pipeline o-- PipelineConfig
|
Pipeline o-- PipelineConfig
|
||||||
Pipeline o-- BaseMaskBuilder
|
Pipeline o-- BaseMaskBuilder
|
||||||
|
Pipeline o-- AutoTokenizer
|
||||||
|
TokenizeTransform o-- AutoTokenizer
|
||||||
|
TokenizeTransform o-- BaseMaskBuilder
|
||||||
|
|
||||||
%% --- Dependency (uses temporarily) ---
|
%% --- Dependency (uses temporarily) ---
|
||||||
TrainConfig ..> BaseStrategy : selects
|
TrainConfig ..> BaseStrategy : selects
|
||||||
@@ -1186,7 +1236,7 @@ classDiagram
|
|||||||
%% --- Association (general usage) ---
|
%% --- Association (general usage) ---
|
||||||
Trainer --> TrainConfig
|
Trainer --> TrainConfig
|
||||||
DPOStrategy --> AutoModel
|
DPOStrategy --> AutoModel
|
||||||
GRPOStrategy --> AutoModel
|
GRPOStrategy --> AutoModel : policy/old/ref
|
||||||
InferenceScheduler --> Task
|
InferenceScheduler --> Task
|
||||||
InferenceScheduler --> TaskStatus
|
InferenceScheduler --> TaskStatus
|
||||||
Task --> TaskStatus
|
Task --> TaskStatus
|
||||||
@@ -1203,8 +1253,8 @@ classDiagram
|
|||||||
| Module | Components | Description |
|
| Module | Components | Description |
|
||||||
|--------|------------|-------------|
|
|--------|------------|-------------|
|
||||||
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
||||||
| **astrai.preprocessing** | BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, Pipeline, filter_by_length, PackingStrategy, PackingStrategyFactory, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory | Declarative JSON-driven data preprocessing |
|
| **astrai.preprocessing** | BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, filter_by_length, PackingStrategy, PackingStrategyFactory, plan_bfd, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory, core (shared helpers) | Declarative JSON-driven data preprocessing |
|
||||||
| **astrai.dataset** | BaseDataset–GRPODataset, Store–JsonlStore/MmapStore/H5Store, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
| **astrai.dataset** | BaseDataset–RecordDataset–DPO/GRPODataset, SEQDataset, SFTDataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlStore, 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 |
|
||||||
@@ -1246,4 +1296,4 @@ classDiagram
|
|||||||
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-07-09
|
> Document Update Time: 2026-07-19
|
||||||
|
|||||||
+36
-18
@@ -61,41 +61,59 @@ StoreFactory.create("bin") → MmapStore
|
|||||||
StoreFactory.create("jsonl") → JsonlStore
|
StoreFactory.create("jsonl") → JsonlStore
|
||||||
```
|
```
|
||||||
|
|
||||||
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage.
|
All three inherit `Store` (base, owns `_data`/`_cum`/`_offsets`/`_normalize`) plus the `Streamable` and `Recordable` mixins, so every backend supports both `fetch(begin, end, keys)` (stream) and `fetch_record(index, keys)` (record) APIs.
|
||||||
|
|
||||||
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`.
|
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. `segments_are_records=True` — each `data_i` dataset is one record.
|
||||||
|
|
||||||
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field.
|
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
|
||||||
|
|
||||||
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based indexing).
|
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field. Two modes: eager (default, applies `TokenizeTransform` to all records at load) and lazy (`processor=fn` given, defers tokenisation to `fetch_record` — used by DPO/GRPO).
|
||||||
|
|
||||||
|
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record-mode indexing). Nested keys (GRPO `responses`/`masks` as `List[List[Tensor]]`) are stored as-is and excluded from both bookkeepings — they are only accessed record-by-record.
|
||||||
|
|
||||||
## Data Keys by Training Type
|
## Data Keys by Training Type
|
||||||
|
|
||||||
| Type | Storage Keys |
|
| Type | Storage Keys | Access Mode |
|
||||||
|------|-------------|
|
|------|-------------|-------------|
|
||||||
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) |
|
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) |
|
||||||
| `sft` | `sequence`, `loss_mask`, `position_ids` |
|
| `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
|
||||||
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` |
|
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
|
||||||
| `grpo` | `prompts`, `responses`, `masks`, `rewards` |
|
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
|
||||||
|
|
||||||
## Dataset Architecture
|
## Dataset Architecture
|
||||||
|
|
||||||
```
|
```
|
||||||
DatasetFactory.load(train_type, load_path, window_size, stride=None, storage_type=None)
|
DatasetFactory.load(train_type, load_path, window_size, stride=None,
|
||||||
|
storage_type=None, tokenizer_path=None,
|
||||||
|
max_len=2048, store=None)
|
||||||
→ BaseDataset.load(load_path, storage_type=None)
|
→ BaseDataset.load(load_path, storage_type=None)
|
||||||
→ detect_format(load_path)
|
→ detect_format(load_path)
|
||||||
→ StoreFactory.create(storage_type)
|
→ StoreFactory.create(storage_type)
|
||||||
→ Store.load(load_path)
|
→ Store.load(load_path)
|
||||||
→ _normalize(raw) # base Store, shared by both backends
|
→ _normalize(raw) # base Store, shared by both backends
|
||||||
→ Store._data[Dict[str, List[Tensor]]] + _cum[Dict[str, List[int]]]
|
→ Store._data[Dict[str, List[Tensor]]]
|
||||||
→ BaseDataset.__getitem__(idx)
|
+ _cum[Dict[str, List[int]]] (stream mode)
|
||||||
→ get_index(idx) → [begin, end)
|
+ _offsets[Dict[str, List[int]]] (record mode)
|
||||||
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
|
|
||||||
|
Stream datasets (SEQ/SFT):
|
||||||
|
BaseDataset.__getitem__(idx)
|
||||||
|
→ get_index(idx) → [begin, end)
|
||||||
|
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
|
||||||
|
|
||||||
|
Record datasets (DPO/GRPO via RecordDataset):
|
||||||
|
RecordDataset.__getitem__(idx)
|
||||||
|
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
|
||||||
```
|
```
|
||||||
|
|
||||||
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). `storage_type` defaults to `None` (auto-detect via `detect_format`).
|
Class hierarchy: `BaseDataset` ← `SEQDataset` / `SFTDataset` (stream); `BaseDataset` ← `RecordDataset` ← `DPODataset` / `GRPODataset` (record).
|
||||||
|
|
||||||
`Store.fetch(begin, end, keys)` accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
|
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
|
||||||
|
|
||||||
|
`tokenizer_path` triggers lazy on-the-fly tokenisation for record datasets on raw JSONL (DPO builds a `dpo_processor`; SEQ/SFT/pre-tokenised backends ignore it). `store` (pre-built `Store`) bypasses `load_path`/`storage_type`/`tokenizer_path` entirely — the caller controls Store construction.
|
||||||
|
|
||||||
|
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
|
||||||
|
|
||||||
|
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present (bin layout with per-record offsets), otherwise indexes `_data[key]` directly (H5/JSONL where each segment is one record).
|
||||||
|
|
||||||
## Sampler
|
## Sampler
|
||||||
|
|
||||||
@@ -109,4 +127,4 @@ DatasetFactory.load(train_type, load_path, window_size, stride=None, storage_typ
|
|||||||
|
|
||||||
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-07-09
|
> Document Update Time: 2026-07-19
|
||||||
|
|||||||
@@ -26,7 +26,7 @@
|
|||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
||||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||||
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
|
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | None |
|
||||||
|
|
||||||
### Optimizer (MuonMix)
|
### Optimizer (MuonMix)
|
||||||
|
|
||||||
@@ -201,4 +201,4 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
> Document Update Time: 2026-07-09
|
> Document Update Time: 2026-07-19
|
||||||
+10
-6
@@ -86,7 +86,7 @@ on_train_end
|
|||||||
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
||||||
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
||||||
|
|
||||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping`.
|
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
|
||||||
|
|
||||||
## Strategies
|
## Strategies
|
||||||
|
|
||||||
@@ -118,7 +118,7 @@ $$
|
|||||||
L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right]
|
L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right]
|
||||||
$$
|
$$
|
||||||
|
|
||||||
Parameters: `beta=0.1`, `reduction="mean"`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`.
|
Parameters: `beta=0.1`, `reduction="sum"`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`.
|
||||||
|
|
||||||
### GRPO (Group Relative Policy Optimization)
|
### GRPO (Group Relative Policy Optimization)
|
||||||
|
|
||||||
@@ -135,10 +135,14 @@ $$
|
|||||||
L_{\text{GRPO}} = -\mathbb{E}_t\left[\min\left(\rho_t A,\; \text{clip}\left(\rho_t, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}_t\left[\frac{\pi_{\text{ref}}}{\pi_\theta} - \log\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\right]
|
L_{\text{GRPO}} = -\mathbb{E}_t\left[\min\left(\rho_t A,\; \text{clip}\left(\rho_t, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}_t\left[\frac{\pi_{\text{ref}}}{\pi_\theta} - \log\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\right]
|
||||||
$$
|
$$
|
||||||
|
|
||||||
where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{ref}}(a_t|s_t)$ is the
|
where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the
|
||||||
per-token probability ratio and the expectations are over valid response tokens.
|
per-token importance sampling ratio against the behaviour policy
|
||||||
|
(`old_model`, synced externally between data-generation rounds) and the
|
||||||
|
expectations are over valid response tokens. The KL term regularises
|
||||||
|
$\pi_\theta$ towards a frozen reference model (`ref_model`, typically
|
||||||
|
the SFT checkpoint).
|
||||||
|
|
||||||
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`, `sync_interval=200`.
|
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`. External sync of `old_model` weights via `sync_old_model()` between data-generation rounds.
|
||||||
|
|
||||||
Keys: `prompts`, `responses`, `masks`, `rewards`.
|
Keys: `prompts`, `responses`, `masks`, `rewards`.
|
||||||
|
|
||||||
@@ -218,4 +222,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-07-09
|
> Document Update Time: 2026-07-19
|
||||||
|
|||||||
+2
-2
@@ -12,7 +12,7 @@ from astrai.config import (
|
|||||||
from astrai.dataset import (
|
from astrai.dataset import (
|
||||||
BaseDataset,
|
BaseDataset,
|
||||||
DatasetFactory,
|
DatasetFactory,
|
||||||
ResumableDistributedSampler,
|
RDSampler,
|
||||||
Store,
|
Store,
|
||||||
StoreFactory,
|
StoreFactory,
|
||||||
)
|
)
|
||||||
@@ -77,7 +77,7 @@ __all__ = [
|
|||||||
"Pipeline",
|
"Pipeline",
|
||||||
"PipelineConfig",
|
"PipelineConfig",
|
||||||
"ProtocolHandler",
|
"ProtocolHandler",
|
||||||
"ResumableDistributedSampler",
|
"RDSampler",
|
||||||
"SamplingPipeline",
|
"SamplingPipeline",
|
||||||
"SchedulerFactory",
|
"SchedulerFactory",
|
||||||
"Store",
|
"Store",
|
||||||
|
|||||||
@@ -37,8 +37,9 @@ class TrainConfig(BaseConfig):
|
|||||||
grad_accum_steps: int = field(
|
grad_accum_steps: int = field(
|
||||||
default=1, metadata={"help": "Number of iterations between steps."}
|
default=1, metadata={"help": "Number of iterations between steps."}
|
||||||
)
|
)
|
||||||
max_grad_norm: float = field(
|
max_grad_norm: Optional[float] = field(
|
||||||
default=1.0, metadata={"help": "Maximum gradient norm."}
|
default=None,
|
||||||
|
metadata={"help": "Maximum gradient norm. None disables clipping."},
|
||||||
)
|
)
|
||||||
gradient_checkpointing_modules: List[str] = field(
|
gradient_checkpointing_modules: List[str] = field(
|
||||||
default_factory=list,
|
default_factory=list,
|
||||||
@@ -87,6 +88,10 @@ class TrainConfig(BaseConfig):
|
|||||||
pin_memory: bool = field(
|
pin_memory: bool = field(
|
||||||
default=False, metadata={"help": "Pin memory for dataloader."}
|
default=False, metadata={"help": "Pin memory for dataloader."}
|
||||||
)
|
)
|
||||||
|
collate_fn: Optional[Callable[[List[Any]], Any]] = field(
|
||||||
|
default=None,
|
||||||
|
metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."},
|
||||||
|
)
|
||||||
|
|
||||||
# distributed training
|
# distributed training
|
||||||
nprocs: int = field(
|
nprocs: int = field(
|
||||||
|
|||||||
@@ -1,15 +1,18 @@
|
|||||||
from astrai.dataset.dataset import (
|
from astrai.dataset.dataset import (
|
||||||
BaseDataset,
|
BaseDataset,
|
||||||
DatasetFactory,
|
DatasetFactory,
|
||||||
|
dpo_collate_fn,
|
||||||
grpo_collate_fn,
|
grpo_collate_fn,
|
||||||
)
|
)
|
||||||
from astrai.dataset.sampler import ResumableDistributedSampler
|
from astrai.dataset.sampler import RDSampler
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
H5Store,
|
H5Store,
|
||||||
JsonlStore,
|
JsonlStore,
|
||||||
MmapStore,
|
MmapStore,
|
||||||
|
Recordable,
|
||||||
Store,
|
Store,
|
||||||
StoreFactory,
|
StoreFactory,
|
||||||
|
Streamable,
|
||||||
detect_format,
|
detect_format,
|
||||||
)
|
)
|
||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
@@ -22,8 +25,11 @@ from astrai.serialization import (
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
"BaseDataset",
|
"BaseDataset",
|
||||||
"DatasetFactory",
|
"DatasetFactory",
|
||||||
|
"dpo_collate_fn",
|
||||||
"grpo_collate_fn",
|
"grpo_collate_fn",
|
||||||
"Store",
|
"Store",
|
||||||
|
"Streamable",
|
||||||
|
"Recordable",
|
||||||
"StoreFactory",
|
"StoreFactory",
|
||||||
"H5Store",
|
"H5Store",
|
||||||
"MmapStore",
|
"MmapStore",
|
||||||
@@ -33,5 +39,5 @@ __all__ = [
|
|||||||
"load_h5",
|
"load_h5",
|
||||||
"save_bin",
|
"save_bin",
|
||||||
"load_bin",
|
"load_bin",
|
||||||
"ResumableDistributedSampler",
|
"RDSampler",
|
||||||
]
|
]
|
||||||
|
|||||||
+349
-232
@@ -1,7 +1,31 @@
|
|||||||
"""Dataset implementations with factory pattern for training."""
|
"""Dataset implementations for training.
|
||||||
|
|
||||||
|
Composition over inheritance — every dataset is a thin wrapper that
|
||||||
|
binds a :class:`Store` to a particular train-type's key mapping. All
|
||||||
|
sample-id → token/record indexing lives on the Store; datasets never
|
||||||
|
know about window/stride math or segment layouts.
|
||||||
|
|
||||||
|
Class hierarchy:
|
||||||
|
|
||||||
|
BaseDataset (ABC) — holds a Store, exposes __len__/keys,
|
||||||
|
overrides __getitem__
|
||||||
|
├── SEQDataset — next-token prediction (stream)
|
||||||
|
├── SFTDataset — loss-mask + position_ids (stream)
|
||||||
|
├── DPODataset — chosen/rejected pairs (record)
|
||||||
|
└── GRPODataset — prompt + response group (record)
|
||||||
|
|
||||||
|
``DatasetFactory.load(train_type, load_path, window_size, stride, …)``
|
||||||
|
builds the Store (auto-detecting format) before constructing the
|
||||||
|
matching dataset. Passing ``store=`` skips Store construction.
|
||||||
|
|
||||||
|
When a record dataset (DPO) reads from raw JSONL, a *processor*
|
||||||
|
function (pure ``record -> Dict[str, Tensor]``) is forwarded to
|
||||||
|
:class:`JsonlStore` so tokenisation happens on the fly.
|
||||||
|
"""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Dict, List, Optional
|
from functools import partial
|
||||||
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
@@ -13,6 +37,147 @@ from astrai.dataset.storage import (
|
|||||||
detect_format,
|
detect_format,
|
||||||
)
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def dpo_tokenize(
|
||||||
|
record: dict,
|
||||||
|
tokenizer,
|
||||||
|
max_len: int = 2048,
|
||||||
|
) -> Optional[dict]:
|
||||||
|
"""Tokenize one DPO record into chosen/rejected + masks.
|
||||||
|
|
||||||
|
Applies the tokenizer's chat template so token sequences match the
|
||||||
|
SFT checkpoint's format. Prompt is rendered with
|
||||||
|
``add_generation_prompt=True``; chosen/rejected are appended as a
|
||||||
|
single assistant turn.
|
||||||
|
|
||||||
|
Accepts:
|
||||||
|
|
||||||
|
- Flat: ``{"prompt": str, "chosen": str, "rejected": str}``
|
||||||
|
- Conv: ``{"prompt": [{role, content}, ...], "chosen": [...], ...}``
|
||||||
|
- Legacy: ``{"input": str, "chosen": str, "rejected": str}``
|
||||||
|
|
||||||
|
No packing, no ``position_ids`` — DPO sequences are independent.
|
||||||
|
"""
|
||||||
|
prompt = record.get("prompt") or record.get("input")
|
||||||
|
chosen = record.get("chosen")
|
||||||
|
rejected = record.get("rejected")
|
||||||
|
if prompt is None or chosen is None or rejected is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
prompt_messages = _to_messages(prompt)
|
||||||
|
chosen_text = _extract_text(chosen)
|
||||||
|
rejected_text = _extract_text(rejected)
|
||||||
|
if chosen_text is None or rejected_text is None:
|
||||||
|
return None
|
||||||
|
chosen_messages = prompt_messages + [{"role": "assistant", "content": chosen_text}]
|
||||||
|
rejected_messages = prompt_messages + [
|
||||||
|
{"role": "assistant", "content": rejected_text}
|
||||||
|
]
|
||||||
|
|
||||||
|
prompt_ids = tokenizer.apply_chat_template(
|
||||||
|
prompt_messages, tokenize=True, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
ch_ids = tokenizer.apply_chat_template(
|
||||||
|
chosen_messages, tokenize=True, add_generation_prompt=False
|
||||||
|
)
|
||||||
|
re_ids = tokenizer.apply_chat_template(
|
||||||
|
rejected_messages, tokenize=True, add_generation_prompt=False
|
||||||
|
)
|
||||||
|
|
||||||
|
full_ch = ch_ids[:max_len]
|
||||||
|
full_re = re_ids[:max_len]
|
||||||
|
|
||||||
|
prompt_len = min(len(prompt_ids), max_len)
|
||||||
|
ch_mask = [0] * prompt_len + [1] * max(0, len(full_ch) - prompt_len)
|
||||||
|
ch_mask = ch_mask[:max_len]
|
||||||
|
re_mask = [0] * prompt_len + [1] * max(0, len(full_re) - prompt_len)
|
||||||
|
re_mask = re_mask[:max_len]
|
||||||
|
|
||||||
|
return {
|
||||||
|
"chosen": full_ch,
|
||||||
|
"rejected": full_re,
|
||||||
|
"chosen_mask": ch_mask,
|
||||||
|
"rejected_mask": re_mask,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _to_messages(value) -> list:
|
||||||
|
"""Accept str or conversation list; return message list."""
|
||||||
|
if isinstance(value, str):
|
||||||
|
return [{"role": "user", "content": value}]
|
||||||
|
if isinstance(value, list):
|
||||||
|
return value
|
||||||
|
return [{"role": "user", "content": str(value)}]
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_text(value) -> Optional[str]:
|
||||||
|
"""Accept str or conversation list; return plain text."""
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
if isinstance(value, str):
|
||||||
|
return value
|
||||||
|
if isinstance(value, list):
|
||||||
|
return "".join(m.get("content", "") for m in value if isinstance(m, dict))
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def dpo_processor(
|
||||||
|
record: dict,
|
||||||
|
tokenizer,
|
||||||
|
max_len: int = 2048,
|
||||||
|
) -> Dict[str, Tensor]:
|
||||||
|
"""DPO processor: wraps :func:`dpo_tokenize` and returns tensors."""
|
||||||
|
result = dpo_tokenize(record, tokenizer, max_len=max_len)
|
||||||
|
if result is None:
|
||||||
|
raise ValueError(f"Malformed DPO record: {list(record.keys())}")
|
||||||
|
return {
|
||||||
|
"chosen": torch.tensor(result["chosen"], dtype=torch.int32),
|
||||||
|
"rejected": torch.tensor(result["rejected"], dtype=torch.int32),
|
||||||
|
"chosen_mask": torch.tensor(result["chosen_mask"], dtype=torch.bool),
|
||||||
|
"rejected_mask": torch.tensor(result["rejected_mask"], dtype=torch.bool),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def dpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||||
|
"""Collate variable-length DPO samples into padded 2-D tensors.
|
||||||
|
|
||||||
|
Input: list of dicts, each with:
|
||||||
|
- chosen: [C_i]
|
||||||
|
- rejected: [R_i]
|
||||||
|
- chosen_mask: [C_i]
|
||||||
|
- rejected_mask: [R_i]
|
||||||
|
|
||||||
|
Output (padded to the max length across chosen/rejected within the batch):
|
||||||
|
- chosen: [B, S_max]
|
||||||
|
- rejected: [B, S_max]
|
||||||
|
- chosen_mask: [B, S_max]
|
||||||
|
- rejected_mask: [B, S_max]
|
||||||
|
"""
|
||||||
|
B = len(batch)
|
||||||
|
S_max = max(b["chosen"].size(0) for b in batch)
|
||||||
|
S_max = max(S_max, max(b["rejected"].size(0) for b in batch))
|
||||||
|
|
||||||
|
chosen = torch.zeros(B, S_max, dtype=torch.long)
|
||||||
|
rejected = torch.zeros(B, S_max, dtype=torch.long)
|
||||||
|
chosen_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||||
|
rejected_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||||
|
|
||||||
|
for i, b in enumerate(batch):
|
||||||
|
c_len = b["chosen"].size(0)
|
||||||
|
r_len = b["rejected"].size(0)
|
||||||
|
chosen[i, :c_len] = b["chosen"]
|
||||||
|
rejected[i, :r_len] = b["rejected"]
|
||||||
|
chosen_mask[i, :c_len] = b["chosen_mask"]
|
||||||
|
rejected_mask[i, :r_len] = b["rejected_mask"]
|
||||||
|
|
||||||
|
return {
|
||||||
|
"chosen": chosen,
|
||||||
|
"rejected": rejected,
|
||||||
|
"chosen_mask": chosen_mask,
|
||||||
|
"rejected_mask": rejected_mask,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||||
@@ -58,200 +223,203 @@ def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
class BaseDataset(Dataset, ABC):
|
def validate_keys(store: Store, required: List[str]) -> None:
|
||||||
"""Abstract base class for all dataset types.
|
"""Raise ``KeyError`` if *store* is missing any *required* key."""
|
||||||
|
if not required:
|
||||||
|
return
|
||||||
|
actual = set(store.keys)
|
||||||
|
missing = [k for k in required if k not in actual]
|
||||||
|
if missing:
|
||||||
|
raise KeyError(
|
||||||
|
f"Store at {getattr(store, '_load_path', '?')} is missing required "
|
||||||
|
f"keys {missing}; available keys are {sorted(actual)}."
|
||||||
|
)
|
||||||
|
|
||||||
Implements common functionality for window-based data fetching.
|
|
||||||
Uses a storage abstraction for format-agnostic data loading.
|
class BaseDataset(Dataset, ABC):
|
||||||
|
"""Abstract base class for dataset types.
|
||||||
|
|
||||||
|
Holds a :class:`Store`. All sample-id indexing is delegated to the
|
||||||
|
store — this class exposes ``__len__`` as ``len(store)`` and the
|
||||||
|
``keys`` property as ``store.keys``. Subclasses implement
|
||||||
|
``__getitem__`` with the train-type-specific key mapping and any
|
||||||
|
training-only index arithmetic (e.g. the next-token ``+1`` shift).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
required_keys: List[str] = []
|
||||||
|
|
||||||
|
def __init__(self, store: Store):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.window_size = window_size
|
self.store: Store = store
|
||||||
self.stride = stride
|
validate_keys(store, self.required_keys)
|
||||||
self.storage: Optional[Store] = None
|
|
||||||
|
|
||||||
@property
|
def __len__(self) -> int:
|
||||||
def required_keys(self) -> List[str]:
|
return len(self.store)
|
||||||
"""Return required storage keys for this dataset type.
|
|
||||||
|
|
||||||
Subclasses should override to specify expected keys.
|
|
||||||
"""
|
|
||||||
return []
|
|
||||||
|
|
||||||
def _validate_keys(self):
|
|
||||||
if not self.required_keys:
|
|
||||||
return
|
|
||||||
actual_keys = set(self.storage.keys)
|
|
||||||
missing = [k for k in self.required_keys if k not in actual_keys]
|
|
||||||
if missing:
|
|
||||||
raise KeyError(
|
|
||||||
f"Dataset {type(self).__name__} requires keys {self.required_keys}, "
|
|
||||||
f"but storage at {self._load_path} only has {sorted(actual_keys)}. "
|
|
||||||
f"Missing: {missing}"
|
|
||||||
)
|
|
||||||
|
|
||||||
def load(self, load_path: str, storage_type: Optional[str] = None, **kwargs):
|
|
||||||
"""Load dataset from the given path.
|
|
||||||
|
|
||||||
Auto-detects the storage format if not specified.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
load_path: Path to the data directory or file
|
|
||||||
storage_type: Force a specific storage type ("h5", "bin", "jsonl"),
|
|
||||||
or None for auto-detection
|
|
||||||
**kwargs: Extra arguments forwarded to the store constructor and
|
|
||||||
to ``store.load()``.
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
KeyError: If the loaded storage is missing required keys.
|
|
||||||
"""
|
|
||||||
if storage_type is None:
|
|
||||||
storage_type = detect_format(load_path)
|
|
||||||
self.storage = StoreFactory.create(storage_type, **kwargs)
|
|
||||||
self._load_path = load_path
|
|
||||||
self.storage.load(load_path, **kwargs)
|
|
||||||
self._validate_keys()
|
|
||||||
|
|
||||||
@property
|
|
||||||
def count(self) -> int:
|
|
||||||
"""Return the total number of raw elements (tokens) in the dataset."""
|
|
||||||
if self.storage is None:
|
|
||||||
return 0
|
|
||||||
return len(self.storage)
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def keys(self) -> List[str]:
|
def keys(self) -> List[str]:
|
||||||
"""Return the available data keys."""
|
return self.store.keys
|
||||||
if self.storage is None:
|
|
||||||
return []
|
|
||||||
return self.storage.keys
|
|
||||||
|
|
||||||
def get_index(self, index: int) -> tuple:
|
@property
|
||||||
"""Calculate begin and end indices for a sample.
|
def token_count(self) -> int:
|
||||||
|
return self.store.token_count
|
||||||
Args:
|
|
||||||
index: Sample index
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (begin_idx, end_idx)
|
|
||||||
"""
|
|
||||||
if self.storage is None:
|
|
||||||
raise RuntimeError("Dataset not loaded, call load() first")
|
|
||||||
total = len(self.storage)
|
|
||||||
if total <= self.window_size:
|
|
||||||
raise ValueError(
|
|
||||||
f"Data too short: {total} tokens <= window_size {self.window_size}"
|
|
||||||
)
|
|
||||||
|
|
||||||
begin_idx = min(index * self.stride, total - 1 - self.window_size)
|
|
||||||
end_idx = min(begin_idx + self.window_size, total - 1)
|
|
||||||
|
|
||||||
return begin_idx, end_idx
|
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
"""Get a single sample by index.
|
|
||||||
|
|
||||||
Must be implemented by subclasses.
|
|
||||||
"""
|
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
if self.storage is None:
|
|
||||||
return 0
|
|
||||||
total = len(self.storage)
|
|
||||||
if total <= self.window_size:
|
|
||||||
return 0
|
|
||||||
return (total - 1 - self.window_size) // self.stride + 1
|
|
||||||
|
|
||||||
|
|
||||||
class DatasetFactory(BaseFactory["BaseDataset"]):
|
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||||
"""Factory class for creating dataset instances.
|
"""Factory for creating dataset instances by train-type.
|
||||||
|
|
||||||
Supports decorator-based registration for extensible dataset types.
|
Use :meth:`DatasetFactory.register("custom")` to register new
|
||||||
All default dataset types (seq, sft, dpo, grpo) are registered automatically
|
dataset classes; they must inherit from :class:`BaseDataset`.
|
||||||
when their classes are defined with the decorator.
|
|
||||||
|
|
||||||
Example usage:
|
|
||||||
@DatasetFactory.register("custom")
|
|
||||||
class CustomDataset(BaseDataset):
|
|
||||||
...
|
|
||||||
|
|
||||||
dataset = DatasetFactory.create("custom", window_size, stride)
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def load(
|
def load(
|
||||||
cls,
|
cls,
|
||||||
train_type: str,
|
train_type: str,
|
||||||
load_path: str,
|
load_path: Optional[str] = None,
|
||||||
window_size: int,
|
window_size: int = 0,
|
||||||
stride: Optional[int] = None,
|
stride: Optional[int] = None,
|
||||||
storage_type: Optional[str] = None,
|
storage_type: Optional[str] = None,
|
||||||
|
tokenizer_path: Optional[str] = None,
|
||||||
|
max_len: int = 2048,
|
||||||
|
store: Optional[Store] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> "BaseDataset":
|
) -> "BaseDataset":
|
||||||
"""Create and load a dataset in one step.
|
"""Create and load a dataset in one step.
|
||||||
|
|
||||||
|
Two entry points:
|
||||||
|
|
||||||
|
- **store given**: bind it directly — the caller fully controls
|
||||||
|
Store construction and processor setup. *load_path*,
|
||||||
|
*storage_type*, *tokenizer_path*, *window_size*, *stride* are
|
||||||
|
ignored.
|
||||||
|
- **store is None**: build a Store from *load_path*, auto-detecting
|
||||||
|
format and constructing a processor when *tokenizer_path* is
|
||||||
|
given for a record dataset on JSONL.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
train_type: Type of training dataset
|
train_type: Registered dataset name ("seq", "sft", "dpo",
|
||||||
load_path: Path to the data file
|
"grpo", …).
|
||||||
window_size: Window size for data sampling
|
load_path: Path to the data file or directory (ignored if
|
||||||
stride: Stride between consecutive samples (default: same as window_size)
|
*store* is given).
|
||||||
storage_type: Storage type ("h5", "bin", "jsonl") or None for auto-detection
|
window_size: Stream window length — only meaningful for
|
||||||
**kwargs: Extra arguments forwarded to ``dataset.load()``.
|
stream datasets (SEQ/SFT). Record datasets ignore it.
|
||||||
|
stride: Stride between consecutive stream samples
|
||||||
|
(default: same as *window_size*).
|
||||||
|
storage_type: Storage backend ("h5", "bin", "jsonl") or
|
||||||
|
None for auto-detection.
|
||||||
|
tokenizer_path: Path to tokenizer for lazy JSONL
|
||||||
|
tokenisation (record datasets only).
|
||||||
|
max_len: Max sequence length forwarded to processors.
|
||||||
|
store: Pre-built, already-loaded Store instance.
|
||||||
|
**kwargs: Extra arguments forwarded to ``store.load()``.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Loaded dataset instance
|
Loaded dataset instance.
|
||||||
"""
|
"""
|
||||||
|
if store is not None:
|
||||||
|
return cls.create(train_type, store=store)
|
||||||
|
|
||||||
|
if load_path is None:
|
||||||
|
raise ValueError("Either load_path or store must be provided")
|
||||||
|
|
||||||
|
if storage_type is None:
|
||||||
|
storage_type = detect_format(load_path)
|
||||||
|
|
||||||
if stride is None:
|
if stride is None:
|
||||||
stride = window_size
|
stride = window_size
|
||||||
|
|
||||||
dataset = cls.create(train_type, window_size, stride)
|
processor = cls._maybe_build_processor(
|
||||||
dataset.load(load_path, storage_type=storage_type, **kwargs)
|
train_type, storage_type, tokenizer_path, max_len
|
||||||
|
)
|
||||||
|
|
||||||
return dataset
|
store_window = cls._store_window_for(train_type, window_size)
|
||||||
|
store = StoreFactory.create(
|
||||||
|
storage_type,
|
||||||
|
window_size=store_window,
|
||||||
|
stride=stride if stride else store_window,
|
||||||
|
)
|
||||||
|
if processor is not None:
|
||||||
|
store.load(load_path, processor=processor, **kwargs)
|
||||||
|
else:
|
||||||
|
store.load(load_path, **kwargs)
|
||||||
|
|
||||||
|
return cls.create(train_type, store=store)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _store_window_for(train_type: str, window_size: int) -> int:
|
||||||
|
"""Stream datasets consume ``window_size``; record datasets ignore it.
|
||||||
|
|
||||||
|
Record datasets (dpo/grpo) treat each record as an independent
|
||||||
|
training unit and never window, so the store is built with
|
||||||
|
``window_size=0`` and ``len(store)`` returns the record count.
|
||||||
|
"""
|
||||||
|
if train_type in ("seq", "sft"):
|
||||||
|
return window_size
|
||||||
|
return 0
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _maybe_build_processor(
|
||||||
|
train_type: str,
|
||||||
|
storage_type: str,
|
||||||
|
tokenizer_path: Optional[str],
|
||||||
|
max_len: int,
|
||||||
|
) -> Optional[Callable[[dict], Dict[str, Tensor]]]:
|
||||||
|
"""Build an on-the-fly tokenisation processor if applicable.
|
||||||
|
|
||||||
|
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
|
||||||
|
pre-tokenised backends (H5/bin) and stream datasets (SEQ/SFT)
|
||||||
|
return ``None`` so no tokenizer is loaded.
|
||||||
|
"""
|
||||||
|
if tokenizer_path is None or storage_type != "jsonl":
|
||||||
|
return None
|
||||||
|
if train_type == "dpo":
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||||
|
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
@DatasetFactory.register("seq")
|
@DatasetFactory.register("seq")
|
||||||
class SEQDataset(BaseDataset):
|
class SEQDataset(BaseDataset):
|
||||||
"""Dataset for sequential next-token prediction training."""
|
"""Dataset for sequential next-token prediction training.
|
||||||
|
|
||||||
@property
|
Stream mode: ``store.fetch(begin, end, "sequence")`` returns the
|
||||||
def required_keys(self) -> List[str]:
|
input window; the +1 shifted call returns the next-token target.
|
||||||
return ["sequence"]
|
"""
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
required_keys = ["sequence"]
|
||||||
return self.storage.fetch(begin_idx, end_idx, "sequence")
|
|
||||||
|
|
||||||
def __getitem__(self, index):
|
def __getitem__(self, index: int):
|
||||||
begin_idx, end_idx = self.get_index(index)
|
begin, end = self.store.sample_window(index)
|
||||||
|
x = self.store.fetch(begin, end, "sequence")
|
||||||
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
|
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||||
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
|
return {
|
||||||
|
"input_ids": x.to(dtype=torch.long),
|
||||||
return {"input_ids": x, "target_ids": y}
|
"target_ids": y.to(dtype=torch.long),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@DatasetFactory.register("sft")
|
@DatasetFactory.register("sft")
|
||||||
class SFTDataset(BaseDataset):
|
class SFTDataset(BaseDataset):
|
||||||
"""Dataset for supervised fine-tuning with loss masking."""
|
"""Dataset for supervised fine-tuning with loss masking.
|
||||||
|
|
||||||
@property
|
Stream mode: ``sequence``/``loss_mask``/``position_ids`` are sliced
|
||||||
def required_keys(self) -> List[str]:
|
to the window. ``loss_mask`` and ``target_ids`` use the +1 shifted
|
||||||
return ["sequence", "loss_mask", "position_ids"]
|
slice so they align with the predicted positions.
|
||||||
|
"""
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
required_keys = ["sequence", "loss_mask", "position_ids"]
|
||||||
return self.storage.fetch(begin_idx, end_idx, key)
|
|
||||||
|
|
||||||
def __getitem__(self, index):
|
|
||||||
begin_idx, end_idx = self.get_index(index)
|
|
||||||
|
|
||||||
x = self._fetch_data(begin_idx, end_idx, "sequence")
|
|
||||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence")
|
|
||||||
position_ids = self._fetch_data(begin_idx, end_idx, "position_ids")
|
|
||||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask")
|
|
||||||
|
|
||||||
|
def __getitem__(self, index: int):
|
||||||
|
begin, end = self.store.sample_window(index)
|
||||||
|
x = self.store.fetch(begin, end, "sequence")
|
||||||
|
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||||
|
position_ids = self.store.fetch(begin, end, "position_ids")
|
||||||
|
loss_mask = self.store.fetch(begin + 1, end + 1, "loss_mask")
|
||||||
return {
|
return {
|
||||||
"input_ids": x.to(dtype=torch.long),
|
"input_ids": x.to(dtype=torch.long),
|
||||||
"target_ids": y.to(dtype=torch.long),
|
"target_ids": y.to(dtype=torch.long),
|
||||||
@@ -262,32 +430,37 @@ class SFTDataset(BaseDataset):
|
|||||||
|
|
||||||
@DatasetFactory.register("dpo")
|
@DatasetFactory.register("dpo")
|
||||||
class DPODataset(BaseDataset):
|
class DPODataset(BaseDataset):
|
||||||
"""Dataset for Direct Preference Optimization training."""
|
"""Record-structured dataset for Direct Preference Optimization.
|
||||||
|
|
||||||
@property
|
Each sample is one preference pair (chosen + rejected) and is an
|
||||||
def required_keys(self) -> List[str]:
|
independent training unit — no windowing, stride, or cross-record
|
||||||
return ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
concatenation. This keeps each sequence self-contained so attention
|
||||||
|
never leaks across preference pairs.
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
Two loading paths (handled by :class:`DatasetFactory`):
|
||||||
return self.storage.fetch(begin_idx, end_idx, key)
|
|
||||||
|
|
||||||
def __getitem__(self, index: int):
|
- **Pre-tokenized** (H5/bin): ``store.load(path)`` reads per-record
|
||||||
begin_idx, end_idx = self.get_index(index)
|
tensors; ``__getitem__`` returns them directly.
|
||||||
|
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
||||||
|
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
||||||
|
no ``position_ids``.
|
||||||
|
"""
|
||||||
|
|
||||||
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
|
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||||
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
|
|
||||||
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
|
|
||||||
dtype=torch.bool
|
|
||||||
)
|
|
||||||
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
|
|
||||||
dtype=torch.bool
|
|
||||||
)
|
|
||||||
|
|
||||||
|
def make_processor(self, tokenizer, max_len: int):
|
||||||
|
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
return {
|
return {
|
||||||
"chosen": chosen,
|
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
|
||||||
"rejected": rejected,
|
"rejected": self.store.fetch_record(index, "rejected").to(dtype=torch.long),
|
||||||
"chosen_mask": chosen_mask,
|
"chosen_mask": self.store.fetch_record(index, "chosen_mask").to(
|
||||||
"rejected_mask": rejected_mask,
|
dtype=torch.bool
|
||||||
|
),
|
||||||
|
"rejected_mask": self.store.fetch_record(index, "rejected_mask").to(
|
||||||
|
dtype=torch.bool
|
||||||
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -295,10 +468,8 @@ class DPODataset(BaseDataset):
|
|||||||
class GRPODataset(BaseDataset):
|
class GRPODataset(BaseDataset):
|
||||||
"""Dataset for offline Group Relative Policy Optimization.
|
"""Dataset for offline Group Relative Policy Optimization.
|
||||||
|
|
||||||
Unlike the window-based datasets (SEQ/SFT/DPO), GRPO data is
|
Each sample is one prompt with its group of responses and scalar
|
||||||
record-structured: each sample is one prompt with its group of
|
rewards — an independent training unit with no windowing or stride.
|
||||||
responses and scalar rewards. There is no windowing or stride —
|
|
||||||
every record is an independent training unit.
|
|
||||||
|
|
||||||
Expected storage layout (produced by JsonlStore or pre-tokenized):
|
Expected storage layout (produced by JsonlStore or pre-tokenized):
|
||||||
|
|
||||||
@@ -308,70 +479,16 @@ class GRPODataset(BaseDataset):
|
|||||||
- ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record
|
- ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, window_size: int = 0, stride: int = 0, **kwargs):
|
required_keys = ["prompts", "responses", "masks", "rewards"]
|
||||||
super().__init__(window_size=window_size, stride=stride or window_size)
|
|
||||||
self._records: List[dict] = []
|
|
||||||
|
|
||||||
@property
|
|
||||||
def required_keys(self) -> List[str]:
|
|
||||||
return ["prompts", "responses", "masks", "rewards"]
|
|
||||||
|
|
||||||
def load(self, load_path: str, storage_type: Optional[str] = None, **kwargs):
|
|
||||||
if storage_type is None:
|
|
||||||
storage_type = detect_format(load_path)
|
|
||||||
self.storage = StoreFactory.create(storage_type, **kwargs)
|
|
||||||
self._load_path = load_path
|
|
||||||
self.storage.load(load_path, **kwargs)
|
|
||||||
self._validate_keys()
|
|
||||||
self._build_records()
|
|
||||||
|
|
||||||
def _validate_keys(self):
|
|
||||||
actual_keys = set(self.storage.keys)
|
|
||||||
missing = [k for k in self.required_keys if k not in actual_keys]
|
|
||||||
if missing:
|
|
||||||
raise KeyError(
|
|
||||||
f"GRPODataset requires keys {self.required_keys}, "
|
|
||||||
f"but storage only has {sorted(actual_keys)}. Missing: {missing}"
|
|
||||||
)
|
|
||||||
|
|
||||||
def _build_records(self):
|
|
||||||
"""Unfold segmented storage into per-record lists.
|
|
||||||
|
|
||||||
``prompts`` is a flat list of 1-D tensors (one per record).
|
|
||||||
``responses`` / ``masks`` are nested lists (G tensors per record).
|
|
||||||
``rewards`` is a flat list of 1-D tensors (len G per record).
|
|
||||||
"""
|
|
||||||
prompt_segs = self.storage._data.get("prompts", [])
|
|
||||||
response_segs = self.storage._data.get("responses", [])
|
|
||||||
mask_segs = self.storage._data.get("masks", [])
|
|
||||||
reward_segs = self.storage._data.get("rewards", [])
|
|
||||||
|
|
||||||
n_records = len(prompt_segs)
|
|
||||||
self._records = []
|
|
||||||
for i in range(n_records):
|
|
||||||
self._records.append(
|
|
||||||
{
|
|
||||||
"prompts": prompt_segs[i],
|
|
||||||
"responses": response_segs[i] if i < len(response_segs) else [],
|
|
||||||
"masks": mask_segs[i] if i < len(mask_segs) else [],
|
|
||||||
"rewards": reward_segs[i]
|
|
||||||
if i < len(reward_segs)
|
|
||||||
else torch.tensor([]),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def count(self) -> int:
|
|
||||||
return len(self._records)
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
return len(self._records)
|
|
||||||
|
|
||||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
rec = self._records[index]
|
prompts = self.store.fetch_record(index, "prompts")
|
||||||
|
responses = self.store.fetch_record(index, "responses")
|
||||||
|
masks = self.store.fetch_record(index, "masks")
|
||||||
|
rewards = self.store.fetch_record(index, "rewards")
|
||||||
return {
|
return {
|
||||||
"prompts": rec["prompts"].to(dtype=torch.long),
|
"prompts": prompts.to(dtype=torch.long),
|
||||||
"responses": [r.to(dtype=torch.long) for r in rec["responses"]],
|
"responses": [r.to(dtype=torch.long) for r in responses],
|
||||||
"masks": [m.to(dtype=torch.bool) for m in rec["masks"]],
|
"masks": [m.to(dtype=torch.bool) for m in masks],
|
||||||
"rewards": rec["rewards"].to(dtype=torch.float32),
|
"rewards": rewards.to(dtype=torch.float32),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,7 +5,15 @@ import torch.distributed as dist
|
|||||||
from torch.utils.data import Dataset, Sampler
|
from torch.utils.data import Dataset, Sampler
|
||||||
|
|
||||||
|
|
||||||
class ResumableDistributedSampler(Sampler[int]):
|
class RDSampler(Sampler[int]):
|
||||||
|
"""Resumable Distributed Sampler.
|
||||||
|
|
||||||
|
A distributed sampler that supports checkpoint-based resume: iteration
|
||||||
|
state (epoch, position) is tracked so training can continue from the
|
||||||
|
exact sample after a restart. Shards the dataset across
|
||||||
|
``dist.world_size`` replicas with optional shuffling.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
data_source: Dataset,
|
data_source: Dataset,
|
||||||
|
|||||||
+481
-190
@@ -1,20 +1,48 @@
|
|||||||
"""Storage backends for different data formats.
|
"""Storage backends for different data formats.
|
||||||
|
|
||||||
Layers:
|
Architecture (composition over inheritance):
|
||||||
- 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:
|
Store (ABC) — owns _data/_cum/_offsets bookkeeping
|
||||||
- Multi-segment: segments kept as-is, no forced concatenation — safe for
|
+ window_size/stride for sample-id
|
||||||
datasets larger than RAM
|
indexing. __getitem__/__len__ produce
|
||||||
- Explicit length: _length = min(total elements across keys), set at load,
|
the smallest iterable unit so Dataset
|
||||||
__len__ returns O(1)
|
classes are pure delegators.
|
||||||
- Zero-copy mmap: MmapStore wraps np.memmap(mode="r"), all DataLoader
|
Streamable (mixin) — raw token slice fetch(begin, end, keys)
|
||||||
workers share OS page-cache pages
|
Recordable (mixin) — raw record slice fetch_record(idx, keys)
|
||||||
|
|
||||||
|
H5Store(Store, Streamable, Recordable)
|
||||||
|
MmapStore(Store, Streamable, Recordable)
|
||||||
|
JsonlStore(Store, Streamable, Recordable)
|
||||||
|
|
||||||
|
Each mixin is a stateless trait that relies on ``self._data`` etc.
|
||||||
|
provided by :class:`Store`. Concrete stores mix in whichever access
|
||||||
|
primitives they support — ``Store`` is the sole base class, so there is
|
||||||
|
no diamond inheritance or MRO ambiguity.
|
||||||
|
|
||||||
|
Sample-id indexing lives on :class:`Store`, not on the dataset:
|
||||||
|
|
||||||
|
- **Stream mode** (``window_size > 0``): ``len(store)`` returns the number
|
||||||
|
of ``(window_size, stride)`` windows that fit in the token river;
|
||||||
|
``store[i]`` returns the *i*-th window as a dict of per-key tensors;
|
||||||
|
``store.sample_window(i)`` exposes the underlying ``(begin, end)``
|
||||||
|
token slice for callers (e.g. next-token trainers) that need a +1
|
||||||
|
shifted companion window.
|
||||||
|
- **Record mode** (``num_records > 0``): ``len(store)`` returns the
|
||||||
|
record count; ``store[i]`` returns the *i*-th record dict.
|
||||||
|
|
||||||
|
Raw token/record access via :meth:`fetch` / :meth:`fetch_record`
|
||||||
|
remains available for low-level callers that want explicit index
|
||||||
|
control. ``store.token_count`` is the total stream token count (what
|
||||||
|
``len(store)`` used to mean in the legacy stream-only API).
|
||||||
|
|
||||||
|
``segments_are_records`` (class attribute on each Store subclass)
|
||||||
|
tells ``_normalize`` whether segments are inherently per-record (H5/
|
||||||
|
JSONL) or opaque shards (bin). Record access for bin relies on
|
||||||
|
``_offsets`` instead.
|
||||||
|
|
||||||
|
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
|
||||||
|
raw records and defers tokenisation to ``fetch_record`` — used by DPO
|
||||||
|
to train directly from a ``.jsonl`` file without a pre-tokenised copy.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import bisect
|
import bisect
|
||||||
@@ -23,20 +51,18 @@ import json
|
|||||||
import logging
|
import logging
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List, Union
|
from typing import Callable, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.preprocess_config import PipelineConfig
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.preprocessing.builder import MaskBuilderFactory
|
from astrai.preprocessing.transform import TokenizeTransform
|
||||||
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
|
||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
load_bin,
|
load_bin,
|
||||||
|
load_bin_offsets,
|
||||||
load_h5,
|
load_h5,
|
||||||
)
|
)
|
||||||
from astrai.tokenize import AutoTokenizer
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -48,7 +74,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", "bin", or "jsonl")
|
Format string ("h5", "bin", "jsonl", or "processed")
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
FileNotFoundError: If no supported data files are found
|
FileNotFoundError: If no supported data files are found
|
||||||
@@ -81,84 +107,173 @@ def detect_format(load_path: str) -> str:
|
|||||||
]
|
]
|
||||||
if jsonl_files:
|
if jsonl_files:
|
||||||
return "jsonl"
|
return "jsonl"
|
||||||
json_files = [
|
|
||||||
Path(p) for p in glob.glob(str(root / "**" / "*.json"), recursive=True)
|
|
||||||
]
|
|
||||||
if json_files:
|
|
||||||
return "jsonl"
|
|
||||||
raise FileNotFoundError(f"No supported data files found at {load_path}")
|
raise FileNotFoundError(f"No supported data files found at {load_path}")
|
||||||
|
|
||||||
|
|
||||||
class Store(ABC):
|
class Store(ABC):
|
||||||
"""String keys -> segmented tensors with ``fetch(begin, end, keys)``.
|
"""Common base for all storage backends.
|
||||||
|
|
||||||
Each key maps to one or more tensor segments (no forced concatenation).
|
A Store owns both its data layout AND its sample-id → token/record
|
||||||
``len(store)`` returns ``self._length`` (explicit, O(1)), the minimum
|
index translation. Datasets are thin wrappers that bind a Store
|
||||||
total element count across all keys.
|
to a particular train-type's key mapping; they never know about
|
||||||
|
window/stride math.
|
||||||
|
|
||||||
Subclasses fill ``self._data`` and ``self._cum`` during ``load()``
|
Two iteration modes:
|
||||||
via ``_normalize()``.
|
|
||||||
|
- **Stream** (``window_size > 0``): data is treated as one long
|
||||||
|
token river. ``len(store)`` returns the number of windows;
|
||||||
|
``store[i]`` slices every stream-compatible key to window ``i``;
|
||||||
|
``store.sample_window(i)`` returns the ``(begin, end)`` token
|
||||||
|
slice for callers needing a +1 shifted companion window.
|
||||||
|
- **Record** (``num_records > 0``): data is per-record.
|
||||||
|
``len(store)`` returns ``num_records``; ``store[i]`` returns
|
||||||
|
the *i*-th record as a dict.
|
||||||
|
|
||||||
|
Raw token slicing is still available via :meth:`fetch` (mixed in
|
||||||
|
by :class:`Streamable`) when a store has stream support configured.
|
||||||
|
Raw record slicing via :meth:`fetch_record` (mixed in by
|
||||||
|
:class:`Recordable`) when a store has record support.
|
||||||
|
|
||||||
|
``token_count`` exposes the raw total stream length — this is what
|
||||||
|
``len(store)`` returned in the legacy stream-only API and what
|
||||||
|
stream-bound ``fetch`` uses for its bounds check.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
segments_are_records: bool = False
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
window_size: int = 0,
|
||||||
|
stride: Optional[int] = None,
|
||||||
|
):
|
||||||
self._data: Dict[str, List[Tensor]] = {}
|
self._data: Dict[str, List[Tensor]] = {}
|
||||||
self._cum: Dict[str, List[int]] = {}
|
self._cum: Dict[str, List[int]] = {}
|
||||||
|
self._offsets: Dict[str, List[int]] = {}
|
||||||
self._length: int = 0
|
self._length: int = 0
|
||||||
|
self._num_records: int = 0
|
||||||
|
self._window_size: int = int(window_size)
|
||||||
|
self._stride: int = int(stride) if stride is not None else int(window_size)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def load(self, path: str) -> None:
|
def load(self, path: str, **kwargs) -> None:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def keys(self) -> List[str]:
|
def keys(self) -> List[str]:
|
||||||
return list(self._data.keys())
|
return list(self._data.keys())
|
||||||
|
|
||||||
def __len__(self) -> int:
|
@property
|
||||||
|
def window_size(self) -> int:
|
||||||
|
return self._window_size
|
||||||
|
|
||||||
|
@property
|
||||||
|
def stride(self) -> int:
|
||||||
|
return self._stride
|
||||||
|
|
||||||
|
@property
|
||||||
|
def token_count(self) -> int:
|
||||||
|
"""Total tokens across all stream segments.
|
||||||
|
|
||||||
|
Useful for the bounds-checked raw :meth:`fetch` and as the
|
||||||
|
legacy ``len(store)`` value.
|
||||||
|
"""
|
||||||
return self._length
|
return self._length
|
||||||
|
|
||||||
def fetch(
|
@property
|
||||||
self,
|
def num_records(self) -> int:
|
||||||
begin: int,
|
"""Number of records available via :meth:`fetch_record`.
|
||||||
end: int,
|
|
||||||
keys: Union[str, List[str]],
|
Non-zero only when the backing layout provides per-record
|
||||||
):
|
indexing (H5/JSONL segments or bin ``_offsets``).
|
||||||
if not self._data:
|
"""
|
||||||
raise RuntimeError("Store not loaded")
|
return self._num_records
|
||||||
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
|
||||||
raise ValueError(
|
@property
|
||||||
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
def num_samples(self) -> int:
|
||||||
|
"""Number of items produced by ``__getitem__``.
|
||||||
|
|
||||||
|
Stream-mode wins when ``window_size > 0`` and there are tokens
|
||||||
|
to slice; otherwise falls back to ``num_records``.
|
||||||
|
"""
|
||||||
|
if self._window_size > 0 and self._length > 0:
|
||||||
|
total = self._length
|
||||||
|
w = self._window_size
|
||||||
|
if total <= w:
|
||||||
|
return 0
|
||||||
|
return (total - 1 - w) // self._stride + 1
|
||||||
|
return self._num_records
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return self.num_samples
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
|
if index < 0:
|
||||||
|
index += self.num_samples
|
||||||
|
if not 0 <= index < self.num_samples:
|
||||||
|
raise IndexError(
|
||||||
|
f"Store index out of range: {index}, num_samples={self.num_samples}"
|
||||||
)
|
)
|
||||||
if isinstance(keys, str):
|
if self._window_size > 0 and self._length > 0:
|
||||||
return self._fetch_key(keys, begin, end)
|
begin, end = self.sample_window(index)
|
||||||
return {k: self._fetch_key(k, begin, end) for k in keys}
|
keys = self._stream_keys()
|
||||||
|
return {k: self.fetch(begin, end, k) for k in keys}
|
||||||
|
return self.fetch_record(index, self._record_keys())
|
||||||
|
|
||||||
def _fetch_key(self, key: str, begin: int, end: int) -> Tensor:
|
def sample_window(self, index: int) -> Tuple[int, int]:
|
||||||
"""Fetch slice [begin, end) across potentially multiple segments."""
|
"""Return ``(begin, end)`` token positions for stream sample *index*.
|
||||||
segments = self._data[key]
|
|
||||||
cum = self._cum[key]
|
|
||||||
seg_start = bisect.bisect_right(cum, begin)
|
|
||||||
seg_end = bisect.bisect_left(cum, end)
|
|
||||||
|
|
||||||
results = []
|
The clipped tail keeps the last reachable window inside the
|
||||||
for i in range(seg_start, seg_end + 1):
|
token river instead of overshooting. Caller is responsible
|
||||||
prev = cum[i - 1] if i > 0 else 0
|
for staying within :attr:`num_samples`: an out-of-range index
|
||||||
s = max(begin - prev, 0)
|
raises ``IndexError``.
|
||||||
e = min(end - prev, segments[i].shape[0])
|
"""
|
||||||
results.append(segments[i][s:e])
|
if self._window_size <= 0:
|
||||||
|
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
|
||||||
|
if self._window_size <= 0 or self._length <= self._window_size:
|
||||||
|
raise IndexError(
|
||||||
|
f"Data too short for window: token_count={self._length}, "
|
||||||
|
f"window_size={self._window_size}"
|
||||||
|
)
|
||||||
|
if not 0 <= index < self.num_samples:
|
||||||
|
raise IndexError(
|
||||||
|
f"Sample index out of range: {index}, num_samples={self.num_samples}"
|
||||||
|
)
|
||||||
|
total = self._length
|
||||||
|
begin = min(index * self._stride, total - 1 - self._window_size)
|
||||||
|
end = min(begin + self._window_size, total - 1)
|
||||||
|
return begin, end
|
||||||
|
|
||||||
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
|
def _stream_keys(self) -> List[str]:
|
||||||
|
out: List[str] = []
|
||||||
|
for k, tensors in self._data.items():
|
||||||
|
if tensors and isinstance(tensors[0], list):
|
||||||
|
continue
|
||||||
|
out.append(k)
|
||||||
|
return out
|
||||||
|
|
||||||
def _normalize(self, raw: Dict[str, list]):
|
def _record_keys(self) -> List[str]:
|
||||||
"""Register segments and pre-compute cumulative lengths.
|
return list(self._data.keys())
|
||||||
|
|
||||||
Does NOT concatenate — segments are kept as-is to avoid OOM on
|
def _normalize(
|
||||||
large datasets. Sets ``self._length`` to the minimum total
|
self,
|
||||||
element count across all flat-tensor keys.
|
raw: Dict[str, list],
|
||||||
|
offsets: Optional[Dict[str, List[int]]] = None,
|
||||||
|
):
|
||||||
|
"""Register segments and pre-compute indices for both access modes.
|
||||||
|
|
||||||
For GRPO multi-response keys, values may be ``List[List[Tensor]]``
|
Stream mode: ``_cum[key]`` accumulates per-segment lengths so
|
||||||
(one list of G tensors per record). These are stored as-is and
|
``Streamable._fetch_stream_key`` can bisect across segments
|
||||||
excluded from the cumulative-length bookkeeping since they are
|
without concatenation.
|
||||||
accessed record-by-record via ``_data`` rather than via ``fetch``.
|
|
||||||
|
Record mode: if *offsets* is provided (bin layout),
|
||||||
|
``_offsets[key]`` stores cumulative per-record offsets into the
|
||||||
|
single concatenated segment. Otherwise, when
|
||||||
|
``segments_are_records`` is True (H5/JSONL), ``_data[key]`` is
|
||||||
|
a per-record list and ``fetch_record`` indexes it directly.
|
||||||
|
|
||||||
|
Nested keys (GRPO ``responses``/``masks`` as
|
||||||
|
``List[List[Tensor]]``) are stored as-is and excluded from both
|
||||||
|
cumulative bookkeepings — they are only accessed record-by-record.
|
||||||
"""
|
"""
|
||||||
flat_lengths = []
|
flat_lengths = []
|
||||||
for key, tensors in raw.items():
|
for key, tensors in raw.items():
|
||||||
@@ -167,7 +282,6 @@ class Store(ABC):
|
|||||||
self._cum[key] = []
|
self._cum[key] = []
|
||||||
flat_lengths.append(0)
|
flat_lengths.append(0)
|
||||||
continue
|
continue
|
||||||
# Skip nested lists (GRPO responses/masks) — record-level access
|
|
||||||
if isinstance(tensors[0], list):
|
if isinstance(tensors[0], list):
|
||||||
self._cum[key] = []
|
self._cum[key] = []
|
||||||
continue
|
continue
|
||||||
@@ -180,166 +294,343 @@ class Store(ABC):
|
|||||||
flat_lengths.append(cum[-1] if cum else 0)
|
flat_lengths.append(cum[-1] if cum else 0)
|
||||||
self._length = min(flat_lengths) if flat_lengths else 0
|
self._length = min(flat_lengths) if flat_lengths else 0
|
||||||
|
|
||||||
|
valid_offsets: Dict[str, List[int]] = {}
|
||||||
|
if offsets:
|
||||||
|
for key, off in offsets.items():
|
||||||
|
segs = self._data.get(key, [])
|
||||||
|
if len(segs) == 1 and len(off) > 1:
|
||||||
|
valid_offsets[key] = off
|
||||||
|
elif len(segs) > 1:
|
||||||
|
logger.warning(
|
||||||
|
"Key '%s' has %d segments with offsets — record mode "
|
||||||
|
"disabled for this key (multi-shard bin+offsets not "
|
||||||
|
"supported). Merge shards or use H5/JSONL.",
|
||||||
|
key,
|
||||||
|
len(segs),
|
||||||
|
)
|
||||||
|
self._offsets = valid_offsets
|
||||||
|
if valid_offsets:
|
||||||
|
record_counts = [len(v) - 1 for v in valid_offsets.values()]
|
||||||
|
self._num_records = min(record_counts) if record_counts else 0
|
||||||
|
elif self.segments_are_records:
|
||||||
|
per_record_counts = []
|
||||||
|
for key, tensors in self._data.items():
|
||||||
|
if tensors and isinstance(tensors[0], list):
|
||||||
|
continue
|
||||||
|
per_record_counts.append(len(tensors))
|
||||||
|
self._num_records = min(per_record_counts) if per_record_counts else 0
|
||||||
|
else:
|
||||||
|
self._num_records = 0
|
||||||
|
|
||||||
|
|
||||||
|
class Streamable:
|
||||||
|
"""Mixin granting raw token-stream access via :meth:`fetch`.
|
||||||
|
|
||||||
|
Stateless trait relying on ``self._data``, ``self._cum``,
|
||||||
|
``self._length`` maintained by :class:`Store`. Stream mode is
|
||||||
|
active when the owning store has ``window_size > 0``; for stores
|
||||||
|
that can also serve record access (H5/JSONL/bin+offsets), the
|
||||||
|
``fetch_record`` API from :class:`Recordable` is used instead.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def fetch(
|
||||||
|
self,
|
||||||
|
begin: int,
|
||||||
|
end: int,
|
||||||
|
keys: Union[str, List[str]],
|
||||||
|
):
|
||||||
|
return _stream_fetch(self, begin, end, keys)
|
||||||
|
|
||||||
|
|
||||||
|
def _stream_fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||||
|
if not getattr(self, "_data", None):
|
||||||
|
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 _fetch_stream_key(self, keys, begin, end)
|
||||||
|
return {k: _fetch_stream_key(self, k, begin, end) for k in keys}
|
||||||
|
|
||||||
|
|
||||||
|
def _fetch_stream_key(self, key: str, begin: int, end: int) -> Tensor:
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
class Recordable:
|
||||||
|
"""Mixin granting raw record access via :meth:`fetch_record`.
|
||||||
|
|
||||||
|
Stateless trait relying on ``self._data``, ``self._offsets``,
|
||||||
|
``self._num_records`` maintained by :class:`Store`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def fetch_record(
|
||||||
|
self,
|
||||||
|
index: int,
|
||||||
|
keys: Union[str, List[str]],
|
||||||
|
):
|
||||||
|
return _record_fetch(self, index, keys)
|
||||||
|
|
||||||
|
|
||||||
|
def _record_fetch(self, index: int, keys: Union[str, List[str]]):
|
||||||
|
if not getattr(self, "_data", None) and self._num_records == 0:
|
||||||
|
raise RuntimeError("Store not loaded")
|
||||||
|
if not 0 <= index < self._num_records:
|
||||||
|
raise ValueError(
|
||||||
|
f"Record index out of bounds: {index}, num_records={self._num_records}"
|
||||||
|
)
|
||||||
|
if isinstance(keys, str):
|
||||||
|
return _fetch_record_key(self, keys, index)
|
||||||
|
return {k: _fetch_record_key(self, k, index) for k in keys}
|
||||||
|
|
||||||
|
|
||||||
|
def _fetch_record_key(self, key: str, index: int):
|
||||||
|
offsets = self._offsets.get(key)
|
||||||
|
if offsets:
|
||||||
|
start = offsets[index]
|
||||||
|
end = (
|
||||||
|
offsets[index + 1]
|
||||||
|
if index + 1 < len(offsets)
|
||||||
|
else self._data[key][0].shape[0]
|
||||||
|
)
|
||||||
|
return self._data[key][0][start:end]
|
||||||
|
return self._data[key][index]
|
||||||
|
|
||||||
|
|
||||||
class StoreFactory(BaseFactory["Store"]):
|
class StoreFactory(BaseFactory["Store"]):
|
||||||
"""Factory for creating Store instances by type name.
|
"""Factory for creating Store instances by type name."""
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
@StoreFactory.register("custom")
|
|
||||||
class CustomStore(Store):
|
|
||||||
...
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
@StoreFactory.register("h5")
|
@StoreFactory.register("h5")
|
||||||
class H5Store(Store):
|
class H5Store(Store, Streamable, Recordable):
|
||||||
"""HDF5-based storage backend (pre-tokenized data)."""
|
"""HDF5-based storage backend (pre-tokenized data).
|
||||||
|
|
||||||
def load(self, path: str):
|
Each key is stored as a group of per-record datasets (``data_0``,
|
||||||
|
``data_1``, …). Supports both access modes:
|
||||||
|
|
||||||
|
- **Stream**: ``fetch(begin, end, key)`` and ``store[i]`` slice
|
||||||
|
across concatenated records via ``_cum`` — used by SEQ/SFT.
|
||||||
|
- **Record**: ``fetch_record(i, key)`` and ``store[i]`` (when
|
||||||
|
``window_size == 0``) index ``_data[key]`` directly — used by
|
||||||
|
DPO/GRPO.
|
||||||
|
"""
|
||||||
|
|
||||||
|
segments_are_records = True
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
window_size: int = 0,
|
||||||
|
stride: Optional[int] = None,
|
||||||
|
):
|
||||||
|
super().__init__(window_size=window_size, stride=stride)
|
||||||
|
|
||||||
|
def load(self, path: str, **kwargs):
|
||||||
self._normalize(load_h5(path))
|
self._normalize(load_h5(path))
|
||||||
|
|
||||||
|
|
||||||
@StoreFactory.register("bin")
|
@StoreFactory.register("bin")
|
||||||
class MmapStore(Store):
|
class MmapStore(Store, Streamable, Recordable):
|
||||||
"""Memory-mapped binary storage backend.
|
"""Memory-mapped binary storage backend.
|
||||||
|
|
||||||
Each key is a single .bin file backed by ``np.memmap(mode="r")``.
|
Each key is a single .bin file backed by ``np.memmap(mode="r")``.
|
||||||
No per-process memory duplication — all DataLoader workers share the
|
No per-process memory duplication — all DataLoader workers share the
|
||||||
same OS page-cache pages.
|
same OS page-cache pages.
|
||||||
|
|
||||||
Format on disk::
|
Supports both access modes:
|
||||||
|
|
||||||
data_root/
|
- **Stream**: always available via :meth:`fetch`.
|
||||||
meta.json # {key: {shape, dtype}, ...}
|
- **Record** (``fetch_record(i, key)``): only when ``meta.json``
|
||||||
<key>.bin # raw numpy array, one per key
|
contains per-record ``offsets`` (written via
|
||||||
|
``save_bin(..., record_keys=...)``). Legacy bin files without
|
||||||
|
offsets have ``num_records == 0`` and ``len(store)`` reflects the
|
||||||
|
windowed sample count when ``window_size > 0``.
|
||||||
|
|
||||||
|
``segments_are_records`` is ``False`` here (bin segments are
|
||||||
|
contiguous streams, not per-record) — record access is driven
|
||||||
|
purely by ``_offsets``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def load(self, path: str):
|
segments_are_records = False
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
window_size: int = 0,
|
||||||
|
stride: Optional[int] = None,
|
||||||
|
):
|
||||||
|
super().__init__(window_size=window_size, stride=stride)
|
||||||
|
self._mmap_refs: List[Tensor] = []
|
||||||
|
|
||||||
|
def load(self, path: str, **kwargs):
|
||||||
self._mmap_refs = []
|
self._mmap_refs = []
|
||||||
root = Path(path)
|
root = Path(path)
|
||||||
all_raw: Dict[str, List[Tensor]] = {}
|
all_raw: Dict[str, List[Tensor]] = {}
|
||||||
|
all_offsets: Dict[str, List[int]] = {}
|
||||||
meta_paths = [
|
meta_paths = [
|
||||||
Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)
|
Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)
|
||||||
]
|
]
|
||||||
for meta_path in meta_paths:
|
for meta_path in meta_paths:
|
||||||
raw = load_bin(str(meta_path.parent))
|
raw = load_bin(str(meta_path.parent))
|
||||||
|
off = load_bin_offsets(str(meta_path.parent))
|
||||||
for key, tensors in raw.items():
|
for key, tensors in raw.items():
|
||||||
if key not in all_raw:
|
if key not in all_raw:
|
||||||
all_raw[key] = []
|
all_raw[key] = []
|
||||||
all_raw[key].extend(tensors)
|
all_raw[key].extend(tensors)
|
||||||
|
for key, o in off.items():
|
||||||
|
if key not in all_offsets:
|
||||||
|
all_offsets[key] = []
|
||||||
|
all_offsets[key].extend(o)
|
||||||
if not meta_paths:
|
if not meta_paths:
|
||||||
raise FileNotFoundError(f"No meta.json found under {path}")
|
raise FileNotFoundError(f"No meta.json found under {path}")
|
||||||
self._normalize(all_raw)
|
self._normalize(all_raw, offsets=all_offsets or None)
|
||||||
for tensors in self._data.values():
|
for tensors in self._data.values():
|
||||||
self._mmap_refs.extend(tensors)
|
self._mmap_refs.extend(tensors)
|
||||||
|
|
||||||
|
|
||||||
@StoreFactory.register("jsonl")
|
class JsonlSource:
|
||||||
class JsonlStore(Store):
|
"""Read raw JSON records from a ``.jsonl`` file or directory.
|
||||||
"""On-the-fly tokenization store for raw JSONL files.
|
|
||||||
|
|
||||||
A JSONL dataset directory contains ``*.jsonl`` files plus a
|
A thin reader used by :class:`JsonlStore` in processor mode — holds
|
||||||
``dataset_config.json`` file that follows the same schema as
|
no tokenizer, performs no tokenisation, just yields dicts.
|
||||||
:class:`PipelineConfig` with an additional ``tokenizer_path`` field.
|
"""
|
||||||
Records are tokenized when the store is loaded and concatenated into
|
|
||||||
segmented tensors matching the key layout expected by the dataset
|
def __init__(self, path: str):
|
||||||
classes (``sequence``, ``loss_mask``, ``position_ids``, ...).
|
self.path = Path(path)
|
||||||
|
self._records: Optional[List[dict]] = None
|
||||||
|
|
||||||
|
def load(self) -> List[dict]:
|
||||||
|
if self._records is None:
|
||||||
|
self._records = self._read(self.path)
|
||||||
|
return self._records
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _read(root: Path) -> List[dict]:
|
||||||
|
if root.is_file():
|
||||||
|
return JsonlSource._read_file(root)
|
||||||
|
return JsonlSource._read_dir(root)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _read_file(path: Path) -> List[dict]:
|
||||||
|
records: List[dict] = []
|
||||||
|
with open(path, "r", encoding="utf-8") as f:
|
||||||
|
for line in f:
|
||||||
|
line = line.strip()
|
||||||
|
if not line:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
records.append(json.loads(line))
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
logger.warning("Failed to parse JSON line in %s, skipping", path)
|
||||||
|
return records
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _read_dir(root: Path) -> List[dict]:
|
||||||
|
records: List[dict] = []
|
||||||
|
for jsonl_path in sorted(root.glob("*.jsonl")):
|
||||||
|
records.extend(JsonlSource._read_file(jsonl_path))
|
||||||
|
return records
|
||||||
|
|
||||||
|
|
||||||
|
@StoreFactory.register("jsonl")
|
||||||
|
class JsonlStore(Store, Streamable, Recordable):
|
||||||
|
"""JSONL reader with two tokenisation modes.
|
||||||
|
|
||||||
|
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
|
||||||
|
files plus (optionally) a ``dataset_config.json`` describing the
|
||||||
|
tokenization pipeline.
|
||||||
|
|
||||||
|
Two modes, selected at :meth:`load` time:
|
||||||
|
|
||||||
|
- **Eager** (default): applies a :class:`TokenizeTransform` to every
|
||||||
|
record at load time and registers per-key tensors via
|
||||||
|
``_normalize``. Both ``fetch`` (stream) and ``fetch_record``
|
||||||
|
(record) work.
|
||||||
|
- **Lazy** (``processor=fn`` passed): keeps raw records and defers
|
||||||
|
tokenisation to ``fetch_record``. Only record access works —
|
||||||
|
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
CONFIG_NAME = "dataset_config.json"
|
CONFIG_NAME = "dataset_config.json"
|
||||||
|
segments_are_records = True
|
||||||
|
|
||||||
def load(self, path: str):
|
def __init__(
|
||||||
root = Path(path)
|
self,
|
||||||
config_path = root / self.CONFIG_NAME
|
window_size: int = 0,
|
||||||
if not config_path.exists():
|
stride: Optional[int] = None,
|
||||||
raise FileNotFoundError(
|
):
|
||||||
f"JSONL dataset config not found: {config_path}. "
|
super().__init__(window_size=window_size, stride=stride)
|
||||||
f"Expected {self.CONFIG_NAME} alongside *.jsonl files."
|
self._source: Optional[JsonlSource] = None
|
||||||
|
self._processor: Optional[Callable[[dict], Dict[str, Tensor]]] = None
|
||||||
|
self._keys_cache: Optional[List[str]] = None
|
||||||
|
|
||||||
|
def load(self, path: str, transform=None, processor=None, **kwargs):
|
||||||
|
self._source = JsonlSource(path)
|
||||||
|
records = self._source.load()
|
||||||
|
|
||||||
|
if processor is not None:
|
||||||
|
self._processor = processor
|
||||||
|
self._num_records = len(records)
|
||||||
|
return
|
||||||
|
|
||||||
|
if transform is None:
|
||||||
|
root = Path(path)
|
||||||
|
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
||||||
|
if config_path is None or not config_path.exists():
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"JSONL dataset config not found. Expected "
|
||||||
|
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
||||||
|
f"explicit transform, or pass processor= for lazy "
|
||||||
|
f"on-the-fly tokenisation."
|
||||||
|
)
|
||||||
|
transform = TokenizeTransform.from_config_file(str(config_path))
|
||||||
|
|
||||||
|
transformed = transform.apply(records)
|
||||||
|
self._normalize(transformed)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def keys(self) -> List[str]:
|
||||||
|
if self._processor is not None:
|
||||||
|
if self._keys_cache is None and self._num_records > 0:
|
||||||
|
sample = self._processor(self._source.load()[0])
|
||||||
|
self._keys_cache = list(sample.keys())
|
||||||
|
return self._keys_cache or []
|
||||||
|
return list(self._data.keys())
|
||||||
|
|
||||||
|
def fetch_record(self, index: int, keys: Union[str, List[str]]):
|
||||||
|
if self._processor is not None:
|
||||||
|
if not 0 <= index < self._num_records:
|
||||||
|
raise ValueError(
|
||||||
|
f"Record index out of bounds: {index}, "
|
||||||
|
f"num_records={self._num_records}"
|
||||||
|
)
|
||||||
|
record = self._source.load()[index]
|
||||||
|
data = self._processor(record)
|
||||||
|
if isinstance(keys, str):
|
||||||
|
return data[keys]
|
||||||
|
return {k: data[k] for k in keys}
|
||||||
|
return _record_fetch(self, index, keys)
|
||||||
|
|
||||||
|
def fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||||
|
if self._processor is not None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"JsonlStore in lazy (processor) mode does not support "
|
||||||
|
"stream fetch(); use fetch_record() instead."
|
||||||
)
|
)
|
||||||
|
return _stream_fetch(self, begin, end, keys)
|
||||||
|
|
||||||
with open(config_path, "r", encoding="utf-8") as f:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
raw_config = json.load(f)
|
if self._processor is not None:
|
||||||
|
return self.fetch_record(index, self._record_keys())
|
||||||
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
|
return super().__getitem__(index)
|
||||||
self.config = PipelineConfig.from_dict(raw_config)
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
|
||||||
mask_builder = MaskBuilderFactory.create("sectioned")
|
|
||||||
position_strategy = PositionIdStrategyFactory.create(
|
|
||||||
self.config.output.position_ids_mode
|
|
||||||
)
|
|
||||||
|
|
||||||
raw: Dict[str, List[Tensor]] = {}
|
|
||||||
doc_sequences: List[List[int]] = []
|
|
||||||
|
|
||||||
def _process_item(item: dict) -> None:
|
|
||||||
nonlocal raw, doc_sequences
|
|
||||||
result = mask_builder.build(item, self.config, tokenizer)
|
|
||||||
if result is None:
|
|
||||||
return
|
|
||||||
result.pop("domain", None)
|
|
||||||
primary_ids = self._primary_ids(result)
|
|
||||||
if not primary_ids:
|
|
||||||
return
|
|
||||||
doc_sequences.append(primary_ids)
|
|
||||||
for key, ids in result.items():
|
|
||||||
if key not in raw:
|
|
||||||
raw[key] = []
|
|
||||||
if ids and isinstance(ids[0], list):
|
|
||||||
# GRPO multi-response: List[List[int]] → List[Tensor]
|
|
||||||
raw[key].append(
|
|
||||||
[torch.tensor(sub, dtype=self._infer_dtype(sub)) for sub in ids]
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
raw[key].append(torch.tensor(ids, dtype=self._infer_dtype(ids)))
|
|
||||||
|
|
||||||
for jsonl_path in sorted(root.glob("*.jsonl")):
|
|
||||||
with open(jsonl_path, "r", encoding="utf-8") as f:
|
|
||||||
for line in f:
|
|
||||||
line = line.strip()
|
|
||||||
if not line:
|
|
||||||
continue
|
|
||||||
try:
|
|
||||||
item = json.loads(line)
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
logger.warning(
|
|
||||||
"Failed to parse JSON line in %s, skipping", jsonl_path
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
_process_item(item)
|
|
||||||
|
|
||||||
for json_path in sorted(root.glob("*.json")):
|
|
||||||
if json_path.name == self.CONFIG_NAME:
|
|
||||||
continue
|
|
||||||
with open(json_path, "r", encoding="utf-8") as f:
|
|
||||||
try:
|
|
||||||
data = json.load(f)
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
logger.warning("Failed to parse JSON file %s, skipping", json_path)
|
|
||||||
continue
|
|
||||||
if isinstance(data, list):
|
|
||||||
for item in data:
|
|
||||||
_process_item(item)
|
|
||||||
elif isinstance(data, dict):
|
|
||||||
_process_item(data)
|
|
||||||
|
|
||||||
pos_ids = position_strategy.generate(doc_sequences)
|
|
||||||
if pos_ids:
|
|
||||||
raw["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
|
||||||
|
|
||||||
self._normalize(raw)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _primary_ids(result: dict) -> List[int]:
|
|
||||||
"""Return the first flat integer list in *result* as the primary id sequence."""
|
|
||||||
for val in result.values():
|
|
||||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
|
||||||
return val
|
|
||||||
return []
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _infer_dtype(ids: List) -> torch.dtype:
|
|
||||||
"""Infer tensor dtype from the first element of a token/value list."""
|
|
||||||
if ids and isinstance(ids[0], float):
|
|
||||||
return torch.float32
|
|
||||||
return torch.int32
|
|
||||||
|
|||||||
@@ -300,7 +300,11 @@ class KVCache(ABC):
|
|||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def bind_tasks(
|
def bind_tasks(
|
||||||
self, task_ids: List[str], total_len: int, device: torch.device
|
self,
|
||||||
|
task_ids: List[str],
|
||||||
|
total_len: int,
|
||||||
|
device: torch.device,
|
||||||
|
write_positions: Optional[Tensor] = None,
|
||||||
) -> CacheView: ...
|
) -> CacheView: ...
|
||||||
|
|
||||||
def task_cached(self, task_id: str) -> int:
|
def task_cached(self, task_id: str) -> int:
|
||||||
@@ -399,7 +403,11 @@ class PageCache(KVCache):
|
|||||||
self._pool.record(page_table[i], prompt_ids, i)
|
self._pool.record(page_table[i], prompt_ids, i)
|
||||||
|
|
||||||
def bind_tasks(
|
def bind_tasks(
|
||||||
self, task_ids: List[str], total_len: int, device: torch.device
|
self,
|
||||||
|
task_ids: List[str],
|
||||||
|
total_len: int,
|
||||||
|
device: torch.device,
|
||||||
|
write_positions: Optional[Tensor] = None,
|
||||||
) -> PageCacheView:
|
) -> PageCacheView:
|
||||||
page_table = self._table.table_tensor(task_ids, device)
|
page_table = self._table.table_tensor(task_ids, device)
|
||||||
return PageCacheView(self._storage, page_table, total_len)
|
return PageCacheView(self._storage, page_table, total_len)
|
||||||
@@ -409,23 +417,37 @@ class ContiguousCacheView(CacheView):
|
|||||||
"""Contiguous KV-cache view for attention layers."""
|
"""Contiguous KV-cache view for attention layers."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, cache: "ContiguousCache", batch_indices: Tensor, total_len: int = 0
|
self,
|
||||||
|
cache: "ContiguousCache",
|
||||||
|
batch_indices: Tensor,
|
||||||
|
total_len: int = 0,
|
||||||
|
write_positions: Optional[Tensor] = None,
|
||||||
):
|
):
|
||||||
self._cache = cache
|
self._cache = cache
|
||||||
self._batch_indices = batch_indices
|
self._batch_indices = batch_indices
|
||||||
self._total_len = total_len
|
self._total_len = total_len
|
||||||
|
self._write_positions = write_positions
|
||||||
|
|
||||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
||||||
seq_len = k.size(1)
|
seq_len = k.size(1)
|
||||||
start_pos = self._total_len - seq_len
|
|
||||||
indices = self._batch_indices
|
indices = self._batch_indices
|
||||||
self._cache.k[layer_id, indices, start_pos : start_pos + seq_len] = k
|
if self._write_positions is not None and seq_len == 1:
|
||||||
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
|
pos = self._write_positions
|
||||||
new_len = start_pos + seq_len
|
self._cache.k[layer_id, indices, pos] = k.squeeze(1)
|
||||||
for s in indices.tolist():
|
self._cache.v[layer_id, indices, pos] = v.squeeze(1)
|
||||||
cur = self._cache._slot_len.get(s, 0)
|
for s, p in zip(indices.tolist(), pos.tolist()):
|
||||||
if new_len > cur:
|
cur = self._cache._slot_len.get(s, 0)
|
||||||
self._cache._slot_len[s] = new_len
|
if p + 1 > cur:
|
||||||
|
self._cache._slot_len[s] = p + 1
|
||||||
|
else:
|
||||||
|
start_pos = self._total_len - seq_len
|
||||||
|
self._cache.k[layer_id, indices, start_pos : start_pos + seq_len] = k
|
||||||
|
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
|
||||||
|
new_len = start_pos + seq_len
|
||||||
|
for s in indices.tolist():
|
||||||
|
cur = self._cache._slot_len.get(s, 0)
|
||||||
|
if new_len > cur:
|
||||||
|
self._cache._slot_len[s] = new_len
|
||||||
|
|
||||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||||
max_len = max(
|
max_len = max(
|
||||||
@@ -491,9 +513,21 @@ class ContiguousCache(KVCache):
|
|||||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||||
return pos < self.max_seq_len
|
return pos < self.max_seq_len
|
||||||
|
|
||||||
|
def task_cached(self, task_id: str) -> int:
|
||||||
|
slot = self._task_slot.get(task_id)
|
||||||
|
if slot is None:
|
||||||
|
return 0
|
||||||
|
return self._slot_len.get(slot, 0)
|
||||||
|
|
||||||
def bind_tasks(
|
def bind_tasks(
|
||||||
self, task_ids: List[str], total_len: int, device: torch.device
|
self,
|
||||||
|
task_ids: List[str],
|
||||||
|
total_len: int,
|
||||||
|
device: torch.device,
|
||||||
|
write_positions: Optional[Tensor] = None,
|
||||||
) -> ContiguousCacheView:
|
) -> ContiguousCacheView:
|
||||||
slots = [self._task_slot[tid] for tid in task_ids]
|
slots = [self._task_slot[tid] for tid in task_ids]
|
||||||
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
|
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
|
||||||
return ContiguousCacheView(self, batch_indices, total_len)
|
return ContiguousCacheView(
|
||||||
|
self, batch_indices, total_len, write_positions=write_positions
|
||||||
|
)
|
||||||
|
|||||||
@@ -106,7 +106,12 @@ class Executor:
|
|||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
outputs = self.model(
|
outputs = self.model(
|
||||||
input_ids.unsqueeze(1),
|
input_ids.unsqueeze(1),
|
||||||
paged_cache=self.kv_cache.bind_tasks(task_ids, total_len, self.device),
|
paged_cache=self.kv_cache.bind_tasks(
|
||||||
|
task_ids,
|
||||||
|
total_len,
|
||||||
|
self.device,
|
||||||
|
write_positions=position_ids,
|
||||||
|
),
|
||||||
position_ids=position_ids.unsqueeze(1),
|
position_ids=position_ids.unsqueeze(1),
|
||||||
)
|
)
|
||||||
logits = outputs["logits"][:, -1, :]
|
logits = outputs["logits"][:, -1, :]
|
||||||
|
|||||||
@@ -138,36 +138,33 @@ class InferenceScheduler:
|
|||||||
t.task_id, t.prompt_ids, start_logical_page
|
t.task_id, t.prompt_ids, start_logical_page
|
||||||
)
|
)
|
||||||
|
|
||||||
pos_groups: Dict[int, List[Task]] = {}
|
decode_tasks = self._task_mgr.get_active_tasks()
|
||||||
for t in self._task_mgr.get_active_tasks():
|
|
||||||
pos_groups.setdefault(t.next_pos, []).append(t)
|
|
||||||
|
|
||||||
for next_pos in sorted(pos_groups.keys()):
|
valid: List[Task] = []
|
||||||
group = sorted(pos_groups[next_pos], key=lambda t: t.task_id)
|
for t in sorted(decode_tasks, key=lambda t: t.task_id):
|
||||||
|
if cache.task_extend(t.task_id, t.next_pos):
|
||||||
|
valid.append(t)
|
||||||
|
else:
|
||||||
|
t.status = TaskStatus.ABORTED
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
|
|
||||||
valid: List[Task] = []
|
if valid:
|
||||||
for t in group:
|
next_tokens = self._executor.execute_decode(valid)
|
||||||
if cache.task_extend(t.task_id, t.next_pos):
|
|
||||||
valid.append(t)
|
for t, ntok in zip(valid, next_tokens):
|
||||||
else:
|
t.output_ids.append(ntok)
|
||||||
t.status = TaskStatus.ABORTED
|
t.output_tokens += 1
|
||||||
|
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||||
|
if new_text:
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||||
|
|
||||||
|
for t in valid:
|
||||||
|
if t.is_finished(stop_ids):
|
||||||
|
remaining = t.flush_remaining(self._task_mgr.tokenizer)
|
||||||
|
if remaining:
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, remaining)
|
||||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
|
|
||||||
if valid:
|
|
||||||
next_tokens = self._executor.execute_decode(valid)
|
|
||||||
|
|
||||||
for t, ntok in zip(valid, next_tokens):
|
|
||||||
t.output_ids.append(ntok)
|
|
||||||
t.output_tokens += 1
|
|
||||||
self._task_mgr.invoke_callback(
|
|
||||||
t.task_id,
|
|
||||||
self._task_mgr.tokenizer.decode([ntok]),
|
|
||||||
)
|
|
||||||
|
|
||||||
for t in valid:
|
|
||||||
if t.is_finished(stop_ids):
|
|
||||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self._stop_event.set()
|
self._stop_event.set()
|
||||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||||
|
|||||||
@@ -13,6 +13,40 @@ logger = logging.getLogger(__name__)
|
|||||||
STOP = object()
|
STOP = object()
|
||||||
|
|
||||||
|
|
||||||
|
class StreamDecoder:
|
||||||
|
"""Incremental decoder for byte-level BPE streaming.
|
||||||
|
|
||||||
|
Byte-level BPE may split a single Unicode character (e.g. em-dash,
|
||||||
|
smart quotes) across multiple tokens. Decoding such a token in
|
||||||
|
isolation produces U+FFFD (replacement char). This decoder
|
||||||
|
accumulates token IDs and only emits text once the trailing
|
||||||
|
characters are complete, buffering incomplete multi-byte sequences
|
||||||
|
until the next token arrives.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("_tokenizer", "_ids", "_emitted")
|
||||||
|
|
||||||
|
def __init__(self, tokenizer: AutoTokenizer):
|
||||||
|
self._tokenizer = tokenizer
|
||||||
|
self._ids: List[int] = []
|
||||||
|
self._emitted: str = ""
|
||||||
|
|
||||||
|
def push(self, token_id: int) -> str:
|
||||||
|
"""Append a token ID and return newly completed text.
|
||||||
|
|
||||||
|
Returns "" while a multi-byte character is still incomplete.
|
||||||
|
"""
|
||||||
|
self._ids.append(token_id)
|
||||||
|
full = self._tokenizer.decode(self._ids, skip_special_tokens=True)
|
||||||
|
if full.endswith("\ufffd"):
|
||||||
|
return ""
|
||||||
|
if len(full) > len(self._emitted):
|
||||||
|
diff = full[len(self._emitted) :]
|
||||||
|
self._emitted = full
|
||||||
|
return diff
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
class TaskStatus(Enum):
|
class TaskStatus(Enum):
|
||||||
"""Task lifecycle states."""
|
"""Task lifecycle states."""
|
||||||
|
|
||||||
@@ -51,6 +85,34 @@ class Task:
|
|||||||
self.output_tokens: int = 0
|
self.output_tokens: int = 0
|
||||||
self.arrival_time = time.time()
|
self.arrival_time = time.time()
|
||||||
self.finish_time: Optional[float] = None
|
self.finish_time: Optional[float] = None
|
||||||
|
self._decoder: Optional[StreamDecoder] = None
|
||||||
|
|
||||||
|
def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
|
||||||
|
"""Decode the last appended output token, buffering incomplete
|
||||||
|
multi-byte sequences across calls.
|
||||||
|
|
||||||
|
Lazily creates a :class:`StreamDecoder` on first use.
|
||||||
|
"""
|
||||||
|
if self._decoder is None:
|
||||||
|
self._decoder = StreamDecoder(tokenizer)
|
||||||
|
return self._decoder.push(self.output_ids[-1])
|
||||||
|
|
||||||
|
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
|
||||||
|
"""Emit any text still buffered in the decoder.
|
||||||
|
|
||||||
|
Called when generation terminates (max_tokens reached, stop
|
||||||
|
sequence, or external removal) to avoid dropping a final
|
||||||
|
incomplete-looking fragment that is actually complete when
|
||||||
|
adjacent to the stop token.
|
||||||
|
"""
|
||||||
|
if self._decoder is None or not self.output_ids:
|
||||||
|
return ""
|
||||||
|
full = tokenizer.decode(self.output_ids, skip_special_tokens=True)
|
||||||
|
if len(full) > len(self._decoder._emitted):
|
||||||
|
diff = full[len(self._decoder._emitted) :]
|
||||||
|
self._decoder._emitted = full
|
||||||
|
return diff
|
||||||
|
return ""
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def next_pos(self) -> int:
|
def next_pos(self) -> int:
|
||||||
|
|||||||
@@ -148,7 +148,14 @@ class BaseExecutor:
|
|||||||
def grad_accum_steps(self) -> int:
|
def grad_accum_steps(self) -> int:
|
||||||
return self.gradient_state.num_steps
|
return self.gradient_state.num_steps
|
||||||
|
|
||||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float:
|
||||||
|
if max_norm is None:
|
||||||
|
total_norm = torch.norm(
|
||||||
|
torch.stack(
|
||||||
|
[p.grad.norm(2) for p in model.parameters() if p.grad is not None]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return total_norm.item()
|
||||||
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
||||||
if isinstance(total_norm, torch.Tensor):
|
if isinstance(total_norm, torch.Tensor):
|
||||||
return total_norm.item()
|
return total_norm.item()
|
||||||
@@ -224,13 +231,6 @@ class DDPExecutor(BaseExecutor):
|
|||||||
return model.module.state_dict()
|
return model.module.state_dict()
|
||||||
return model.state_dict()
|
return model.state_dict()
|
||||||
|
|
||||||
def _gather_state_dict(self, model: nn.Module):
|
|
||||||
if not self.use_distributed:
|
|
||||||
return self.unwrap_model(model)
|
|
||||||
if get_rank() != 0:
|
|
||||||
return None
|
|
||||||
return self.unwrap_model(model)
|
|
||||||
|
|
||||||
|
|
||||||
@ExecutorFactory.register("fsdp")
|
@ExecutorFactory.register("fsdp")
|
||||||
class FSDPExecutor(BaseExecutor):
|
class FSDPExecutor(BaseExecutor):
|
||||||
@@ -289,7 +289,9 @@ class FSDPExecutor(BaseExecutor):
|
|||||||
return model.no_sync()
|
return model.no_sync()
|
||||||
return contextlib.nullcontext()
|
return contextlib.nullcontext()
|
||||||
|
|
||||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float:
|
||||||
|
if max_norm is None:
|
||||||
|
return super().clip_grad_norm(model, max_norm)
|
||||||
if isinstance(model, FSDP) and self.use_distributed:
|
if isinstance(model, FSDP) and self.use_distributed:
|
||||||
total_norm = model.clip_grad_norm_(max_norm)
|
total_norm = model.clip_grad_norm_(max_norm)
|
||||||
if isinstance(total_norm, torch.Tensor):
|
if isinstance(total_norm, torch.Tensor):
|
||||||
|
|||||||
@@ -8,12 +8,14 @@ from astrai.preprocessing.builder import (
|
|||||||
from astrai.preprocessing.packing import (
|
from astrai.preprocessing.packing import (
|
||||||
PackingStrategy,
|
PackingStrategy,
|
||||||
PackingStrategyFactory,
|
PackingStrategyFactory,
|
||||||
|
plan_bfd,
|
||||||
)
|
)
|
||||||
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
|
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
|
||||||
from astrai.preprocessing.position_id import (
|
from astrai.preprocessing.position_id import (
|
||||||
PositionIdStrategy,
|
PositionIdStrategy,
|
||||||
PositionIdStrategyFactory,
|
PositionIdStrategyFactory,
|
||||||
)
|
)
|
||||||
|
from astrai.preprocessing.transform import TokenizeTransform
|
||||||
from astrai.preprocessing.writer import (
|
from astrai.preprocessing.writer import (
|
||||||
StoreWriter,
|
StoreWriter,
|
||||||
StoreWriterFactory,
|
StoreWriterFactory,
|
||||||
@@ -32,5 +34,7 @@ __all__ = [
|
|||||||
"SingleOutputMaskBuilder",
|
"SingleOutputMaskBuilder",
|
||||||
"StoreWriter",
|
"StoreWriter",
|
||||||
"StoreWriterFactory",
|
"StoreWriterFactory",
|
||||||
|
"TokenizeTransform",
|
||||||
"filter_by_length",
|
"filter_by_length",
|
||||||
|
"plan_bfd",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,124 @@
|
|||||||
|
"""Shared preprocessing kernel used by both :class:`Pipeline` and
|
||||||
|
:class:`TokenizeTransform`.
|
||||||
|
|
||||||
|
The two entry points previously duplicated ~60 % of their logic:
|
||||||
|
record iteration, mask-builder invocation, primary-id extraction,
|
||||||
|
per-key accumulation, dtype inference and position-id generation.
|
||||||
|
This module factors out the common core as pure functions so that
|
||||||
|
the online (``TokenizeTransform``) and offline (``Pipeline``) paths
|
||||||
|
stay in lockstep.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from itertools import chain
|
||||||
|
from typing import Dict, Iterator, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
|
from astrai.preprocessing.builder import MaskBuilderFactory
|
||||||
|
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def build_preprocessing_components(config: PipelineConfig, tokenizer_path: str):
|
||||||
|
"""Load tokenizer, mask builder and position-id strategy together.
|
||||||
|
|
||||||
|
Both ``Pipeline`` and ``TokenizeTransform`` need the same triple;
|
||||||
|
centralising the construction avoids drift (e.g. one path forgetting
|
||||||
|
to create the position-id strategy).
|
||||||
|
"""
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||||
|
mask_builder = MaskBuilderFactory.create("sectioned")
|
||||||
|
position_strategy = PositionIdStrategyFactory.create(
|
||||||
|
config.output.position_ids_mode
|
||||||
|
)
|
||||||
|
return tokenizer, mask_builder, position_strategy
|
||||||
|
|
||||||
|
|
||||||
|
def primary_ids(result: dict) -> List[int]:
|
||||||
|
"""Return the first flat int-list value in *result*.
|
||||||
|
|
||||||
|
Used for token counting and position-id generation when the
|
||||||
|
primary key name is not known (DPO uses ``chosen``, GRPO uses
|
||||||
|
``prompts``, SFT uses ``sequence``).
|
||||||
|
"""
|
||||||
|
for val in result.values():
|
||||||
|
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||||
|
return val
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def infer_dtype(ids: List) -> torch.dtype:
|
||||||
|
"""Float values become float32, everything else int32."""
|
||||||
|
if ids and isinstance(ids[0], float):
|
||||||
|
return torch.float32
|
||||||
|
return torch.int32
|
||||||
|
|
||||||
|
|
||||||
|
def iter_raw_records(
|
||||||
|
records: List[dict],
|
||||||
|
mask_builder,
|
||||||
|
config: PipelineConfig,
|
||||||
|
tokenizer,
|
||||||
|
) -> Iterator[dict]:
|
||||||
|
"""Yield mask-builder output dicts for each record, skipping failures.
|
||||||
|
|
||||||
|
Drops ``domain`` from the result (callers that need it should read
|
||||||
|
it before calling this). Each yielded dict maps a key
|
||||||
|
(``sequence``, ``chosen``, ``responses``…) to either a flat
|
||||||
|
``List[int]`` or a nested ``List[List[int]]`` (GRPO responses/masks).
|
||||||
|
"""
|
||||||
|
for item in records:
|
||||||
|
result = mask_builder.build(item, config, tokenizer)
|
||||||
|
if result is None:
|
||||||
|
continue
|
||||||
|
result.pop("domain", None)
|
||||||
|
if not primary_ids(result):
|
||||||
|
continue
|
||||||
|
yield result
|
||||||
|
|
||||||
|
|
||||||
|
def to_per_record_tensors(
|
||||||
|
raw: Dict[str, list],
|
||||||
|
) -> Dict[str, List[torch.Tensor]]:
|
||||||
|
"""Convert an accumulated ``{key: [per-record ids]}`` dict to tensors.
|
||||||
|
|
||||||
|
Handles three shapes transparently:
|
||||||
|
|
||||||
|
- ``List[int]`` per record (``sequence``, ``chosen``…) → one tensor per record.
|
||||||
|
- ``List[List[int]]`` per record (GRPO ``responses``/``masks``) → one
|
||||||
|
``List[Tensor]`` per record (nested), preserving the per-response
|
||||||
|
boundary so downstream code can index responses individually.
|
||||||
|
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||||
|
|
||||||
|
The detection mirrors the previous inline logic in
|
||||||
|
``Pipeline._flush`` and ``TokenizeTransform.apply``.
|
||||||
|
"""
|
||||||
|
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||||
|
for key, ids_list in raw.items():
|
||||||
|
if ids_list and isinstance(ids_list[0], list):
|
||||||
|
tensors[key] = [
|
||||||
|
[torch.tensor(sub, dtype=infer_dtype(sub)) for sub in ids]
|
||||||
|
if ids and isinstance(ids[0], list)
|
||||||
|
else torch.tensor(ids, dtype=infer_dtype(ids))
|
||||||
|
for ids in ids_list
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
tensors[key] = [
|
||||||
|
torch.tensor(list(chain.from_iterable(ids_list)), dtype=torch.int32)
|
||||||
|
]
|
||||||
|
return tensors
|
||||||
|
|
||||||
|
|
||||||
|
def build_position_ids(
|
||||||
|
sequences: List[List[int]],
|
||||||
|
strategy,
|
||||||
|
) -> Optional[List[int]]:
|
||||||
|
"""Generate position ids for *sequences* using *strategy*.
|
||||||
|
|
||||||
|
Returns ``None`` when the strategy produces no ids (e.g. ``none``
|
||||||
|
mode), so callers can skip attaching the key instead of storing
|
||||||
|
an empty list.
|
||||||
|
"""
|
||||||
|
pos_ids = strategy.generate(sequences)
|
||||||
|
return pos_ids or None
|
||||||
@@ -19,6 +19,43 @@ def _truncate(seq: List[int], max_len: int, mode: str) -> List[int]:
|
|||||||
return seq[:max_len]
|
return seq[:max_len]
|
||||||
|
|
||||||
|
|
||||||
|
def plan_bfd(
|
||||||
|
sequences: List[List[int]], max_packed_len: int, truncation_mode: str = "keep_start"
|
||||||
|
) -> List[List[int]]:
|
||||||
|
"""Best-Fit Decreasing bin packing of *sequences* into bins.
|
||||||
|
|
||||||
|
Returns a list of bins, each bin a list of original indices into
|
||||||
|
*sequences*. Bin capacities are respected on the *truncated*
|
||||||
|
length of each sequence (so a sequence longer than
|
||||||
|
*max_packed_len* counts at *max_packed_len*).
|
||||||
|
|
||||||
|
Pure index-based so callers can apply the same plan to any
|
||||||
|
aligned key (``loss_mask``, ``position_ids``…).
|
||||||
|
"""
|
||||||
|
n = len(sequences)
|
||||||
|
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
|
||||||
|
bins: List[List[int]] = []
|
||||||
|
bin_lengths: List[int] = []
|
||||||
|
|
||||||
|
for orig_idx in order:
|
||||||
|
seq_len = len(_truncate(sequences[orig_idx], max_packed_len, truncation_mode))
|
||||||
|
best_bin = None
|
||||||
|
best_remain = max_packed_len + 1
|
||||||
|
for i, bl in enumerate(bin_lengths):
|
||||||
|
remain = max_packed_len - bl
|
||||||
|
if seq_len <= remain < best_remain:
|
||||||
|
best_remain = remain
|
||||||
|
best_bin = i
|
||||||
|
if best_bin is not None:
|
||||||
|
bins[best_bin].append(orig_idx)
|
||||||
|
bin_lengths[best_bin] += seq_len
|
||||||
|
else:
|
||||||
|
bins.append([orig_idx])
|
||||||
|
bin_lengths.append(seq_len)
|
||||||
|
|
||||||
|
return bins
|
||||||
|
|
||||||
|
|
||||||
class PackingStrategy(ABC):
|
class PackingStrategy(ABC):
|
||||||
"""Reorder and truncate sequences within a shard."""
|
"""Reorder and truncate sequences within a shard."""
|
||||||
|
|
||||||
@@ -70,7 +107,7 @@ class BFDPacking(PackingStrategy):
|
|||||||
sequences = keys.get("sequence", [])
|
sequences = keys.get("sequence", [])
|
||||||
if not sequences:
|
if not sequences:
|
||||||
return keys
|
return keys
|
||||||
bins = self._plan(sequences, max_packed_len, truncation_mode)
|
bins = plan_bfd(sequences, max_packed_len, truncation_mode)
|
||||||
|
|
||||||
packed: Dict[str, List[List[int]]] = {}
|
packed: Dict[str, List[List[int]]] = {}
|
||||||
for k, vals in keys.items():
|
for k, vals in keys.items():
|
||||||
@@ -91,35 +128,6 @@ class BFDPacking(PackingStrategy):
|
|||||||
result.extend(vals[i])
|
result.extend(vals[i])
|
||||||
return result
|
return result
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _plan(
|
|
||||||
sequences: List[List[int]], max_packed_len: int, truncation_mode: str
|
|
||||||
) -> List[List[int]]:
|
|
||||||
n = len(sequences)
|
|
||||||
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
|
|
||||||
bins: List[List[int]] = []
|
|
||||||
bin_lengths: List[int] = []
|
|
||||||
|
|
||||||
for orig_idx in order:
|
|
||||||
seq_len = len(
|
|
||||||
_truncate(sequences[orig_idx], max_packed_len, truncation_mode)
|
|
||||||
)
|
|
||||||
best_bin = None
|
|
||||||
best_remain = max_packed_len + 1
|
|
||||||
for i, bl in enumerate(bin_lengths):
|
|
||||||
remain = max_packed_len - bl
|
|
||||||
if seq_len <= remain < best_remain:
|
|
||||||
best_remain = remain
|
|
||||||
best_bin = i
|
|
||||||
if best_bin is not None:
|
|
||||||
bins[best_bin].append(orig_idx)
|
|
||||||
bin_lengths[best_bin] += seq_len
|
|
||||||
else:
|
|
||||||
bins.append([orig_idx])
|
|
||||||
bin_lengths.append(seq_len)
|
|
||||||
|
|
||||||
return bins
|
|
||||||
|
|
||||||
|
|
||||||
@PackingStrategyFactory.register("bfd_split")
|
@PackingStrategyFactory.register("bfd_split")
|
||||||
class BFDSplitPacking(BFDPacking):
|
class BFDSplitPacking(BFDPacking):
|
||||||
|
|||||||
@@ -4,6 +4,10 @@ Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
|
|||||||
sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id
|
sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id
|
||||||
generation and storage writing are each delegated to pluggable strategies,
|
generation and storage writing are each delegated to pluggable strategies,
|
||||||
dispatched by configuration keys.
|
dispatched by configuration keys.
|
||||||
|
|
||||||
|
Record iteration, mask building, primary-id extraction and per-key
|
||||||
|
accumulation are shared with :class:`TokenizeTransform` via the
|
||||||
|
:mod:`astrai.preprocessing.core` helpers.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
@@ -17,11 +21,13 @@ import torch
|
|||||||
import tqdm
|
import tqdm
|
||||||
|
|
||||||
from astrai.config.preprocess_config import PipelineConfig
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
from astrai.preprocessing.builder import MaskBuilderFactory
|
from astrai.preprocessing.core import (
|
||||||
|
build_preprocessing_components,
|
||||||
|
iter_raw_records,
|
||||||
|
primary_ids,
|
||||||
|
)
|
||||||
from astrai.preprocessing.packing import PackingStrategyFactory
|
from astrai.preprocessing.packing import PackingStrategyFactory
|
||||||
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
|
||||||
from astrai.preprocessing.writer import StoreWriterFactory
|
from astrai.preprocessing.writer import StoreWriterFactory
|
||||||
from astrai.tokenize import AutoTokenizer
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -64,20 +70,18 @@ class Pipeline:
|
|||||||
self.output_dir = output_dir
|
self.output_dir = output_dir
|
||||||
self.tokenizer_path = tokenizer_path
|
self.tokenizer_path = tokenizer_path
|
||||||
|
|
||||||
self.mask_builder = MaskBuilderFactory.create("sectioned")
|
self.tokenizer, self.mask_builder, self._position_id = (
|
||||||
|
build_preprocessing_components(config, tokenizer_path)
|
||||||
|
)
|
||||||
self._packer = PackingStrategyFactory.create(
|
self._packer = PackingStrategyFactory.create(
|
||||||
config.preprocessing.packing_strategy
|
config.preprocessing.packing_strategy
|
||||||
)
|
)
|
||||||
self._position_id = PositionIdStrategyFactory.create(
|
|
||||||
config.output.position_ids_mode
|
|
||||||
)
|
|
||||||
self._writer = StoreWriterFactory.create(config.output.storage_format)
|
self._writer = StoreWriterFactory.create(config.output.storage_format)
|
||||||
|
|
||||||
def transform(self, item: dict) -> Optional[dict]:
|
def transform(self, item: dict) -> Optional[dict]:
|
||||||
return self.mask_builder.build(item, self.config, self._tokenizer)
|
return self.mask_builder.build(item, self.config, self.tokenizer)
|
||||||
|
|
||||||
def run(self):
|
def run(self):
|
||||||
self._tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_path)
|
|
||||||
domains: dict = defaultdict(lambda: defaultdict(list))
|
domains: dict = defaultdict(lambda: defaultdict(list))
|
||||||
total_tokens = 0
|
total_tokens = 0
|
||||||
shard_idx: dict[str, int] = defaultdict(int)
|
shard_idx: dict[str, int] = defaultdict(int)
|
||||||
@@ -102,14 +106,7 @@ class Pipeline:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
domain = result.pop("domain", "__default__")
|
domain = result.pop("domain", "__default__")
|
||||||
|
ids = primary_ids(result)
|
||||||
is_multi = bool(getattr(self.config.input, "sources", None))
|
|
||||||
if is_multi:
|
|
||||||
ids = self._primary_ids(result)
|
|
||||||
else:
|
|
||||||
ids = result.pop("sequence")
|
|
||||||
result["sequence"] = ids
|
|
||||||
|
|
||||||
if not ids:
|
if not ids:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -129,15 +126,6 @@ class Pipeline:
|
|||||||
if total_tokens > 0:
|
if total_tokens > 0:
|
||||||
self._flush(domains, shard_idx)
|
self._flush(domains, shard_idx)
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _primary_ids(result: dict) -> list:
|
|
||||||
"""Return the first list-valued entry in *result* as the primary id
|
|
||||||
sequence for token counting."""
|
|
||||||
for val in result.values():
|
|
||||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
|
||||||
return val
|
|
||||||
return []
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _align_bucket(bucket: dict, result: dict, ids: list):
|
def _align_bucket(bucket: dict, result: dict, ids: list):
|
||||||
"""Pad previously-accumulated keys that are missing from *result*."""
|
"""Pad previously-accumulated keys that are missing from *result*."""
|
||||||
@@ -170,39 +158,12 @@ class Pipeline:
|
|||||||
original_sequences = keys.get("sequence", [])
|
original_sequences = keys.get("sequence", [])
|
||||||
mode = self.config.output.position_ids_mode
|
mode = self.config.output.position_ids_mode
|
||||||
|
|
||||||
if mode == "doc_reset" and original_sequences:
|
keys = self._inject_doc_reset_position_ids(keys, mode, original_sequences)
|
||||||
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
|
|
||||||
|
|
||||||
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
||||||
|
tensors = self._to_tensors(keys)
|
||||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
tensors = self._inject_continuous_position_ids(
|
||||||
for key, ids_list in keys.items():
|
tensors, mode, keys.get("sequence", [])
|
||||||
dt = _STR_TO_DTYPE.get(
|
)
|
||||||
self.config.output.dtype.get(key, "int32"), torch.int32
|
|
||||||
)
|
|
||||||
# GRPO multi-response keys store List[List[int]] per record
|
|
||||||
# (responses/masks). Rewards store List[float] per record.
|
|
||||||
# Both produce List[Tensor] (one tensor per record), but
|
|
||||||
# responses need inner flattening while rewards do not.
|
|
||||||
if ids_list and isinstance(ids_list[0], list):
|
|
||||||
tensors[key] = [
|
|
||||||
torch.tensor(
|
|
||||||
list(chain.from_iterable(ids))
|
|
||||||
if ids and isinstance(ids[0], list)
|
|
||||||
else ids,
|
|
||||||
dtype=dt,
|
|
||||||
)
|
|
||||||
for ids in ids_list
|
|
||||||
]
|
|
||||||
else:
|
|
||||||
tensors[key] = [
|
|
||||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
|
||||||
]
|
|
||||||
|
|
||||||
if mode == "continuous" and original_sequences:
|
|
||||||
pos_ids = self._position_id.generate(keys.get("sequence", []))
|
|
||||||
if pos_ids:
|
|
||||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
|
||||||
|
|
||||||
self._writer.save(self.output_dir, domain, idx, tensors)
|
self._writer.save(self.output_dir, domain, idx, tensors)
|
||||||
shard_idx[domain] = idx + 1
|
shard_idx[domain] = idx + 1
|
||||||
@@ -212,3 +173,76 @@ class Pipeline:
|
|||||||
f" saved {domain}/shard_{idx:04d} "
|
f" saved {domain}/shard_{idx:04d} "
|
||||||
f"({tensors[first_key][0].numel():,} tokens)"
|
f"({tensors[first_key][0].numel():,} tokens)"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _inject_doc_reset_position_ids(
|
||||||
|
self,
|
||||||
|
keys: Dict[str, list],
|
||||||
|
mode: str,
|
||||||
|
original_sequences: List[List[int]],
|
||||||
|
) -> Dict[str, list]:
|
||||||
|
"""Attach per-document position_ids before packing (``doc_reset``).
|
||||||
|
|
||||||
|
``doc_reset`` position ids must enter the packer so that each
|
||||||
|
packed bin concatenates the per-doc ranges in bin order. The
|
||||||
|
per-record structure ``[range(len(s)) for s in seqs]`` is required
|
||||||
|
by the packer (it concatenates per-record lists per bin); the
|
||||||
|
``PositionIdStrategy.generate`` flattens, so it cannot be used
|
||||||
|
directly here — it is only consulted for the ``continuous``
|
||||||
|
post-packing path.
|
||||||
|
"""
|
||||||
|
if mode != "doc_reset" or not original_sequences:
|
||||||
|
return keys
|
||||||
|
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
|
||||||
|
return keys
|
||||||
|
|
||||||
|
def _inject_continuous_position_ids(
|
||||||
|
self,
|
||||||
|
tensors: Dict[str, List[torch.Tensor]],
|
||||||
|
mode: str,
|
||||||
|
packed_sequences: List[List[int]],
|
||||||
|
) -> Dict[str, List[torch.Tensor]]:
|
||||||
|
"""Attach a single continuous position_ids tensor after packing.
|
||||||
|
|
||||||
|
``continuous`` mode spans the whole shard (post-packing), so it
|
||||||
|
cannot participate in bin packing — it is computed from the
|
||||||
|
packed sequences and appended directly to the tensor dict.
|
||||||
|
"""
|
||||||
|
if mode != "continuous" or not packed_sequences:
|
||||||
|
return tensors
|
||||||
|
pos_ids = self._position_id.generate(packed_sequences)
|
||||||
|
if pos_ids:
|
||||||
|
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||||
|
return tensors
|
||||||
|
|
||||||
|
def _to_tensors(self, keys: Dict[str, list]) -> Dict[str, List[torch.Tensor]]:
|
||||||
|
"""Convert packed per-key id lists to tensors.
|
||||||
|
|
||||||
|
Honours ``config.output.dtype`` overrides per key; falls back to
|
||||||
|
``int32``. Handles three shapes (see
|
||||||
|
:func:`astrai.preprocessing.core.to_per_record_tensors` for the
|
||||||
|
equivalent online-path helper):
|
||||||
|
- ``List[int]`` per record → one tensor per record.
|
||||||
|
- ``List[List[int]]`` per record (GRPO responses/masks) → one tensor
|
||||||
|
per record, inner lists flattened.
|
||||||
|
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||||
|
"""
|
||||||
|
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||||
|
for key, ids_list in keys.items():
|
||||||
|
dt = _STR_TO_DTYPE.get(
|
||||||
|
self.config.output.dtype.get(key, "int32"), torch.int32
|
||||||
|
)
|
||||||
|
if ids_list and isinstance(ids_list[0], list):
|
||||||
|
tensors[key] = [
|
||||||
|
torch.tensor(
|
||||||
|
list(chain.from_iterable(ids))
|
||||||
|
if ids and isinstance(ids[0], list)
|
||||||
|
else ids,
|
||||||
|
dtype=dt,
|
||||||
|
)
|
||||||
|
for ids in ids_list
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
tensors[key] = [
|
||||||
|
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
||||||
|
]
|
||||||
|
return tensors
|
||||||
|
|||||||
@@ -0,0 +1,92 @@
|
|||||||
|
"""Tokenization transform for JSONL record streams.
|
||||||
|
|
||||||
|
Bridges the Reader layer (``JsonlStore`` reads raw JSON records) and the
|
||||||
|
Dataset layer (expects per-record tensors). Holds the tokenizer,
|
||||||
|
mask-builder and position-id strategy together so that I/O code stays
|
||||||
|
free of model dependencies.
|
||||||
|
|
||||||
|
The record-processing core (mask building, primary-id extraction,
|
||||||
|
per-key tensorisation, position-id generation) is shared with
|
||||||
|
:class:`astrai.preprocessing.pipeline.Pipeline` via the
|
||||||
|
:mod:`astrai.preprocessing.core` helpers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
|
from astrai.preprocessing.core import (
|
||||||
|
build_position_ids,
|
||||||
|
build_preprocessing_components,
|
||||||
|
iter_raw_records,
|
||||||
|
to_per_record_tensors,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TokenizeTransform:
|
||||||
|
"""Tokenize raw JSONL record dicts into per-key tensor lists.
|
||||||
|
|
||||||
|
Owns the three preprocessing concerns that were previously inlined in
|
||||||
|
``JsonlStore``: tokenization, loss-mask construction and position-id
|
||||||
|
generation. Constructing it loads the tokenizer, so it is intentionally
|
||||||
|
cheap to pass around once built.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: Pipeline config describing sections / masks / position mode.
|
||||||
|
tokenizer_path: Path passed to ``AutoTokenizer.from_pretrained``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, config: PipelineConfig, tokenizer_path: str):
|
||||||
|
self.config = config
|
||||||
|
self.tokenizer, self.mask_builder, self.position_strategy = (
|
||||||
|
build_preprocessing_components(config, tokenizer_path)
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_config_file(cls, config_path: str) -> "TokenizeTransform":
|
||||||
|
"""Build from a ``dataset_config.json`` file path.
|
||||||
|
|
||||||
|
The config file follows :class:`PipelineConfig` schema with an
|
||||||
|
extra ``tokenizer_path`` field. When omitted, the config's
|
||||||
|
parent directory is used as the tokenizer path.
|
||||||
|
"""
|
||||||
|
root = Path(config_path).parent
|
||||||
|
with open(config_path, "r", encoding="utf-8") as f:
|
||||||
|
raw_config = json.load(f)
|
||||||
|
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
|
||||||
|
config = PipelineConfig.from_dict(raw_config)
|
||||||
|
return cls(config, tokenizer_path)
|
||||||
|
|
||||||
|
def apply(self, records: List[dict]) -> Dict[str, list]:
|
||||||
|
"""Tokenize a list of raw record dicts.
|
||||||
|
|
||||||
|
Returns a dict mapping key (``sequence``, ``chosen``, ``responses``,
|
||||||
|
…) to a list of per-record tensors (or nested tensor lists for
|
||||||
|
multi-response keys such as GRPO ``responses``).
|
||||||
|
"""
|
||||||
|
raw: Dict[str, list] = {}
|
||||||
|
doc_sequences: List[List[int]] = []
|
||||||
|
|
||||||
|
for result in iter_raw_records(
|
||||||
|
records, self.mask_builder, self.config, self.tokenizer
|
||||||
|
):
|
||||||
|
primary = None
|
||||||
|
for val in result.values():
|
||||||
|
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||||
|
primary = val
|
||||||
|
break
|
||||||
|
if primary is not None:
|
||||||
|
doc_sequences.append(primary)
|
||||||
|
for key, ids in result.items():
|
||||||
|
raw.setdefault(key, []).append(ids)
|
||||||
|
|
||||||
|
tensors = to_per_record_tensors(raw)
|
||||||
|
|
||||||
|
pos_ids = build_position_ids(doc_sequences, self.position_strategy)
|
||||||
|
if pos_ids is not None:
|
||||||
|
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||||
|
|
||||||
|
return tensors
|
||||||
@@ -19,6 +19,7 @@ from astrai.serialization.checkpoint import (
|
|||||||
)
|
)
|
||||||
from astrai.serialization.dataset import (
|
from astrai.serialization.dataset import (
|
||||||
load_bin,
|
load_bin,
|
||||||
|
load_bin_offsets,
|
||||||
load_h5,
|
load_h5,
|
||||||
save_bin,
|
save_bin,
|
||||||
save_h5,
|
save_h5,
|
||||||
@@ -37,6 +38,7 @@ __all__ = [
|
|||||||
"save_safetensors",
|
"save_safetensors",
|
||||||
"save_torch",
|
"save_torch",
|
||||||
"load_bin",
|
"load_bin",
|
||||||
|
"load_bin_offsets",
|
||||||
"load_h5",
|
"load_h5",
|
||||||
"save_bin",
|
"save_bin",
|
||||||
"save_h5",
|
"save_h5",
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
import h5py
|
import h5py
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -50,12 +50,43 @@ def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
|||||||
return tensor_group
|
return tensor_group
|
||||||
|
|
||||||
|
|
||||||
def save_bin(file_path: str, tensor_group: Dict[str, List[Tensor]]):
|
def save_bin(
|
||||||
|
file_path: str,
|
||||||
|
tensor_group: Dict[str, List[Tensor]],
|
||||||
|
record_keys: Optional[List[str]] = None,
|
||||||
|
):
|
||||||
|
"""Save tensors as memory-mapped binary files.
|
||||||
|
|
||||||
|
When *record_keys* is provided, those keys are written with per-record
|
||||||
|
cumulative offsets in ``meta.json`` so that ``MmapStore.fetch_record``
|
||||||
|
can slice individual records from the concatenated binary without
|
||||||
|
cross-record concatenation. Keys not in *record_keys* (e.g. SEQ
|
||||||
|
``sequence``) are written as a single contiguous stream without
|
||||||
|
offsets, preserving backward compatibility.
|
||||||
|
|
||||||
|
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
||||||
|
not supported in bin format — use H5 for those.
|
||||||
|
"""
|
||||||
os.makedirs(file_path, exist_ok=True)
|
os.makedirs(file_path, exist_ok=True)
|
||||||
|
record_keys = set(record_keys or [])
|
||||||
meta = {}
|
meta = {}
|
||||||
for key, tensors in tensor_group.items():
|
for key, tensors in tensor_group.items():
|
||||||
|
if tensors and isinstance(tensors[0], list):
|
||||||
|
raise ValueError(
|
||||||
|
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
||||||
|
f"in bin format. Use H5 or JSONL storage instead."
|
||||||
|
)
|
||||||
cat = torch.cat(tensors, dim=0)
|
cat = torch.cat(tensors, dim=0)
|
||||||
meta[key] = {"shape": list(cat.shape), "dtype": str(cat.dtype).split(".")[-1]}
|
entry: Dict[str, Any] = {
|
||||||
|
"shape": list(cat.shape),
|
||||||
|
"dtype": str(cat.dtype).split(".")[-1],
|
||||||
|
}
|
||||||
|
if key in record_keys:
|
||||||
|
offsets = [0]
|
||||||
|
for t in tensors:
|
||||||
|
offsets.append(offsets[-1] + t.shape[0])
|
||||||
|
entry["offsets"] = offsets
|
||||||
|
meta[key] = entry
|
||||||
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
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:
|
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
||||||
json.dump(meta, f)
|
json.dump(meta, f)
|
||||||
@@ -74,3 +105,19 @@ def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
|||||||
)
|
)
|
||||||
segments[key] = [torch.from_numpy(arr)]
|
segments[key] = [torch.from_numpy(arr)]
|
||||||
return segments
|
return segments
|
||||||
|
|
||||||
|
|
||||||
|
def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
|
||||||
|
"""Read per-record cumulative offsets from ``meta.json``.
|
||||||
|
|
||||||
|
Returns an empty dict when no key has offsets (legacy bin files),
|
||||||
|
in which case record-mode access falls back to per-record segment
|
||||||
|
indexing (H5/JSONL layout).
|
||||||
|
"""
|
||||||
|
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||||
|
meta = json.load(f)
|
||||||
|
offsets: Dict[str, List[int]] = {}
|
||||||
|
for key, info in meta.items():
|
||||||
|
if "offsets" in info:
|
||||||
|
offsets[key] = info["offsets"]
|
||||||
|
return offsets
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
from functools import cached_property
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
from jinja2 import Template
|
from jinja2 import Template
|
||||||
@@ -29,7 +30,19 @@ class ChatTemplate:
|
|||||||
self.description = description
|
self.description = description
|
||||||
self.default_variables = default_variables or {}
|
self.default_variables = default_variables or {}
|
||||||
self.special_tokens = special_tokens or {}
|
self.special_tokens = special_tokens or {}
|
||||||
self._compiled: Template = Template(template_str)
|
|
||||||
|
@cached_property
|
||||||
|
def _compiled(self) -> Template:
|
||||||
|
"""Lazy-compiled Jinja2 template, cached on first access.
|
||||||
|
|
||||||
|
The compiled :class:`~jinja2.Template` holds a dynamically-generated
|
||||||
|
``root`` render function whose ``__module__`` is ``None``; under
|
||||||
|
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
|
||||||
|
multiprocessing. By deferring compilation to first access, the
|
||||||
|
default pickle protocol serialises only ``template_str``; each
|
||||||
|
worker rebuilds the cache on first render.
|
||||||
|
"""
|
||||||
|
return Template(self.template_str)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_string(
|
def from_string(
|
||||||
|
|||||||
@@ -164,7 +164,14 @@ class AutoTokenizer:
|
|||||||
- tokenizer.bos_token → returns string
|
- tokenizer.bos_token → returns string
|
||||||
- tokenizer.bos_token_id → returns corresponding integer ID
|
- tokenizer.bos_token_id → returns corresponding integer ID
|
||||||
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
||||||
|
|
||||||
|
Internal/private attrs are not intercepted: during unpickle
|
||||||
|
``__dict__`` is empty, so probing ``self._special_token_map``
|
||||||
|
would recurse infinitely.
|
||||||
"""
|
"""
|
||||||
|
if key.startswith("_"):
|
||||||
|
raise AttributeError(key)
|
||||||
|
|
||||||
# Handle stop_ids - return IDs for all special tokens
|
# Handle stop_ids - return IDs for all special tokens
|
||||||
if key == "stop_ids":
|
if key == "stop_ids":
|
||||||
stop_ids = []
|
stop_ids = []
|
||||||
|
|||||||
@@ -98,7 +98,6 @@ class BaseStrategy(ABC):
|
|||||||
self.model = model
|
self.model = model
|
||||||
self.device = device
|
self.device = device
|
||||||
self.executor = kwargs.pop("executor", None)
|
self.executor = kwargs.pop("executor", None)
|
||||||
self.model_fn = kwargs.pop("model_fn", None)
|
|
||||||
self.extra_kwargs = kwargs
|
self.extra_kwargs = kwargs
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
@@ -225,7 +224,7 @@ class DPOStrategy(BaseStrategy):
|
|||||||
device: str,
|
device: str,
|
||||||
ref_model: nn.Module,
|
ref_model: nn.Module,
|
||||||
beta: float = 0.1,
|
beta: float = 0.1,
|
||||||
reduction: str = "mean",
|
reduction: str = "sum",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(model, device, **kwargs)
|
super().__init__(model, device, **kwargs)
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import torch.nn as nn
|
|||||||
from torch.utils.data import DataLoader, random_split
|
from torch.utils.data import DataLoader, random_split
|
||||||
|
|
||||||
from astrai.config.train_config import TrainConfig
|
from astrai.config.train_config import TrainConfig
|
||||||
from astrai.dataset import ResumableDistributedSampler
|
from astrai.dataset import RDSampler
|
||||||
from astrai.model.components.lora import inject_lora
|
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
|
||||||
@@ -141,7 +141,7 @@ class TrainContextBuilder:
|
|||||||
)
|
)
|
||||||
|
|
||||||
sampler_offset = context.consumed_samples // context.world_size
|
sampler_offset = context.consumed_samples // context.world_size
|
||||||
sampler = ResumableDistributedSampler(
|
sampler = RDSampler(
|
||||||
data_source=train_dataset,
|
data_source=train_dataset,
|
||||||
start_epoch=context.epoch,
|
start_epoch=context.epoch,
|
||||||
start_iter=sampler_offset,
|
start_iter=sampler_offset,
|
||||||
@@ -154,10 +154,11 @@ class TrainContextBuilder:
|
|||||||
num_workers=cfg.num_workers,
|
num_workers=cfg.num_workers,
|
||||||
pin_memory=cfg.pin_memory,
|
pin_memory=cfg.pin_memory,
|
||||||
prefetch_factor=cfg.prefetch_factor,
|
prefetch_factor=cfg.prefetch_factor,
|
||||||
|
collate_fn=cfg.collate_fn,
|
||||||
)
|
)
|
||||||
|
|
||||||
if val_dataset is not None:
|
if val_dataset is not None:
|
||||||
val_sampler = ResumableDistributedSampler(
|
val_sampler = RDSampler(
|
||||||
data_source=val_dataset,
|
data_source=val_dataset,
|
||||||
start_epoch=0,
|
start_epoch=0,
|
||||||
start_iter=0,
|
start_iter=0,
|
||||||
@@ -171,6 +172,7 @@ class TrainContextBuilder:
|
|||||||
num_workers=cfg.num_workers,
|
num_workers=cfg.num_workers,
|
||||||
pin_memory=cfg.pin_memory,
|
pin_memory=cfg.pin_memory,
|
||||||
prefetch_factor=cfg.prefetch_factor,
|
prefetch_factor=cfg.prefetch_factor,
|
||||||
|
collate_fn=cfg.collate_fn,
|
||||||
)
|
)
|
||||||
|
|
||||||
context.model, context.optimizer, context.dataloader, context.scheduler = (
|
context.model, context.optimizer, context.dataloader, context.scheduler = (
|
||||||
@@ -209,7 +211,6 @@ class TrainContextBuilder:
|
|||||||
model=context.model,
|
model=context.model,
|
||||||
device=device,
|
device=device,
|
||||||
executor=executor,
|
executor=executor,
|
||||||
model_fn=cfg.model_fn,
|
|
||||||
**strategy_kwargs,
|
**strategy_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from huggingface_hub import snapshot_download
|
|||||||
|
|
||||||
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||||
DEFAULT_LOCAL_DIR = Path(PROJECT_ROOT, "params")
|
DEFAULT_LOCAL_DIR = Path(PROJECT_ROOT, "params")
|
||||||
DEFAULT_REPO_ID = "ViperEk/KHAOSZ"
|
DEFAULT_REPO_ID = "ViperEkura/AstrAI-V1-instruct"
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
|
|||||||
@@ -26,11 +26,9 @@ def batch_generate():
|
|||||||
|
|
||||||
prompts = [
|
prompts = [
|
||||||
tokenizer.apply_chat_template(
|
tokenizer.apply_chat_template(
|
||||||
[
|
[{"role": "user", "content": q}],
|
||||||
{"role": "system", "content": "You are a helpful assistant."},
|
|
||||||
{"role": "user", "content": q},
|
|
||||||
],
|
|
||||||
tokenize=False,
|
tokenize=False,
|
||||||
|
add_generation_prompt=True,
|
||||||
)
|
)
|
||||||
for q in inputs
|
for q in inputs
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -58,8 +58,8 @@ def parse_args():
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--system_prompt",
|
"--system_prompt",
|
||||||
type=str,
|
type=str,
|
||||||
default="You are a helpful assistant.",
|
default="",
|
||||||
help="Optional system prompt",
|
help="Optional system prompt (default: empty, model not SFT-trained on system role)",
|
||||||
)
|
)
|
||||||
return parser.parse_args()
|
return parser.parse_args()
|
||||||
|
|
||||||
@@ -73,18 +73,20 @@ def chat():
|
|||||||
model.to(device="cuda", dtype=torch.bfloat16)
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
engine = InferenceEngine(model=model, tokenizer=tokenizer)
|
engine = InferenceEngine(model=model, tokenizer=tokenizer)
|
||||||
|
|
||||||
messages = [{"role": "system", "content": args.system_prompt}]
|
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
query = input(">> ")
|
query = input(">> ")
|
||||||
if query == "!exit":
|
if query == "!exit":
|
||||||
break
|
break
|
||||||
|
|
||||||
messages.append({"role": "user", "content": query})
|
msgs = []
|
||||||
|
if args.system_prompt:
|
||||||
|
msgs.append({"role": "system", "content": args.system_prompt})
|
||||||
|
msgs.append({"role": "user", "content": query})
|
||||||
|
prompt = tokenizer.apply_chat_template(
|
||||||
|
msgs, tokenize=False, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
|
||||||
full_response = ""
|
full_response = ""
|
||||||
prompt = tokenizer.apply_chat_template(messages, tokenize=False)
|
|
||||||
|
|
||||||
for token in engine.generate(
|
for token in engine.generate(
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
stream=True,
|
stream=True,
|
||||||
@@ -99,7 +101,6 @@ def chat():
|
|||||||
full_response += token
|
full_response += token
|
||||||
|
|
||||||
print()
|
print()
|
||||||
messages.append({"role": "assistant", "content": full_response.strip()})
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from typing import Dict, Iterator, List, Optional, Sequence, Tuple
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
import tqdm
|
import tqdm
|
||||||
|
from datasets import load_dataset
|
||||||
|
|
||||||
from astrai.inference import InferenceEngine
|
from astrai.inference import InferenceEngine
|
||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
@@ -29,9 +30,7 @@ from astrai.tokenize import AutoTokenizer
|
|||||||
# Config
|
# Config
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
HUMANEVAL_URL = (
|
HUMANEVAL_HF_DATASET = "openai/openai_humaneval"
|
||||||
"https://github.com/openai/human-eval/raw/master/data/HumanEval.jsonl.gz"
|
|
||||||
)
|
|
||||||
|
|
||||||
STOP_SEQUENCES = [
|
STOP_SEQUENCES = [
|
||||||
"\nclass ",
|
"\nclass ",
|
||||||
@@ -64,21 +63,16 @@ class EvalConfig:
|
|||||||
problem_indices: Optional[List[int]] = None
|
problem_indices: Optional[List[int]] = None
|
||||||
|
|
||||||
|
|
||||||
def download(url: str, path: str):
|
def download(path: str):
|
||||||
if os.path.exists(path):
|
if os.path.exists(path):
|
||||||
return
|
return
|
||||||
import gzip
|
|
||||||
import urllib.request
|
|
||||||
|
|
||||||
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||||
print(f"Downloading {url} ...")
|
print(f"Downloading HumanEval from HuggingFace ({HUMANEVAL_HF_DATASET}) ...")
|
||||||
tmp = path + ".tmp"
|
ds = load_dataset(HUMANEVAL_HF_DATASET, split="test")
|
||||||
urllib.request.urlretrieve(url, tmp)
|
with open(path, "w", encoding="utf-8") as f:
|
||||||
with gzip.open(tmp, "rb") as f_in:
|
for item in ds:
|
||||||
with open(path, "wb") as f_out:
|
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||||
f_out.write(f_in.read())
|
print(f" saved {len(ds)} problems to {path}")
|
||||||
os.remove(tmp)
|
|
||||||
print(f" saved to {path}")
|
|
||||||
|
|
||||||
|
|
||||||
def load_jsonl(path: str) -> List[dict]:
|
def load_jsonl(path: str) -> List[dict]:
|
||||||
@@ -318,7 +312,7 @@ def run_pipeline(cfg: EvalConfig) -> Dict:
|
|||||||
with open(cfg.test_only, encoding="utf-8") as f:
|
with open(cfg.test_only, encoding="utf-8") as f:
|
||||||
generated = json.load(f)
|
generated = json.load(f)
|
||||||
else:
|
else:
|
||||||
download(HUMANEVAL_URL, cfg.data_path)
|
download(cfg.data_path)
|
||||||
|
|
||||||
problems = load_jsonl(cfg.data_path)
|
problems = load_jsonl(cfg.data_path)
|
||||||
if cfg.problem_indices:
|
if cfg.problem_indices:
|
||||||
|
|||||||
@@ -26,28 +26,22 @@ import torch.nn.functional as F
|
|||||||
import tqdm
|
import tqdm
|
||||||
|
|
||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
|
from astrai.preprocessing.packing import plan_bfd
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
def _pack_bins(pairs, max_len):
|
def _pack_bins(pairs, max_len):
|
||||||
"""BFD bin packing: pack (c+r) into bins of max total length."""
|
"""BFD bin packing: pack (c+r) into bins of max total length.
|
||||||
indexed = sorted(enumerate(pairs), key=lambda x: -(len(x[1][0]) + len(x[1][1])))
|
|
||||||
bins = []
|
Reuses :func:`plan_bfd` so the BFD heuristic stays single-sourced.
|
||||||
lengths = []
|
"""
|
||||||
for orig_idx, (c, r) in indexed:
|
# Treat each pair as a single sequence of length len(c)+len(r) for
|
||||||
size = len(c) + len(r)
|
# planning purposes; plan_bfd works on pure lengths.
|
||||||
best_bin = -1
|
fake_sequences = [[0] * (len(c) + len(r)) for c, r in pairs]
|
||||||
for bi, rem in enumerate(lengths):
|
plan = plan_bfd(fake_sequences, max_len)
|
||||||
if rem >= size:
|
return [
|
||||||
if best_bin < 0 or rem < lengths[best_bin]:
|
[(i, pairs[i][0], pairs[i][1]) for i in bin_indices] for bin_indices in plan
|
||||||
best_bin = bi
|
]
|
||||||
if best_bin >= 0:
|
|
||||||
bins[best_bin].append((orig_idx, c, r))
|
|
||||||
lengths[best_bin] -= size
|
|
||||||
else:
|
|
||||||
bins.append([(orig_idx, c, r)])
|
|
||||||
lengths.append(max_len - size)
|
|
||||||
return bins
|
|
||||||
|
|
||||||
|
|
||||||
def _resolve_sentinel_ids(tokenizer, sentinel_text):
|
def _resolve_sentinel_ids(tokenizer, sentinel_text):
|
||||||
|
|||||||
@@ -14,21 +14,17 @@ import argparse
|
|||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
import urllib.request
|
|
||||||
from typing import Callable, Dict, List, Optional
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import tqdm
|
import tqdm
|
||||||
|
from datasets import load_dataset
|
||||||
|
|
||||||
from astrai.inference import InferenceEngine
|
from astrai.inference import InferenceEngine
|
||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
IFEVAL_URL = (
|
IFEVAL_HF_DATASET = "google/IFEval"
|
||||||
"https://raw.githubusercontent.com/google-research/"
|
|
||||||
"google-research/master/instruction_following_eval/data/input_data.jsonl"
|
|
||||||
)
|
|
||||||
|
|
||||||
CONSTRAINT_VERIFIERS: Dict[str, Callable[[str, dict], bool]] = {}
|
CONSTRAINT_VERIFIERS: Dict[str, Callable[[str, dict], bool]] = {}
|
||||||
|
|
||||||
|
|
||||||
@@ -310,15 +306,12 @@ def download_ifeval(data_path: str):
|
|||||||
if os.path.exists(data_path):
|
if os.path.exists(data_path):
|
||||||
return
|
return
|
||||||
os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True)
|
os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True)
|
||||||
print(f"Downloading IFEval from {IFEVAL_URL} ...")
|
print(f"Downloading IFEval from HuggingFace ({IFEVAL_HF_DATASET}) ...")
|
||||||
tmp = data_path + ".tmp"
|
ds = load_dataset(IFEVAL_HF_DATASET, split="train")
|
||||||
urllib.request.urlretrieve(IFEVAL_URL, tmp)
|
with open(data_path, "w", encoding="utf-8") as f:
|
||||||
with open(tmp, "rb") as f_in:
|
for item in ds:
|
||||||
content = f_in.read()
|
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||||
with open(data_path, "wb") as f_out:
|
print(f" saved {len(ds)} items to {data_path}")
|
||||||
f_out.write(content)
|
|
||||||
os.remove(tmp)
|
|
||||||
print(f" saved to {data_path}")
|
|
||||||
|
|
||||||
|
|
||||||
def load_problems(data_path: str) -> List[dict]:
|
def load_problems(data_path: str) -> List[dict]:
|
||||||
|
|||||||
@@ -4,18 +4,18 @@ import argparse
|
|||||||
import csv
|
import csv
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import shutil
|
import random
|
||||||
import tarfile
|
from collections import defaultdict
|
||||||
|
|
||||||
import requests
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
import tqdm
|
import tqdm
|
||||||
|
from datasets import load_dataset
|
||||||
|
|
||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
MMLU_URL = "https://people.eecs.berkeley.edu/~hendrycks/data.tar"
|
MMLU_HF_DATASET = "cais/mmlu"
|
||||||
MMLU_SUBJECTS = [
|
MMLU_SUBJECTS = [
|
||||||
"abstract_algebra",
|
"abstract_algebra",
|
||||||
"anatomy",
|
"anatomy",
|
||||||
@@ -77,38 +77,40 @@ MMLU_SUBJECTS = [
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
def _download_and_extract(url: str, data_dir: str):
|
def _write_subject_csv(data_dir: str, split: str, subject: str, rows: list[dict]):
|
||||||
tar_path = os.path.join(data_dir, "data.tar")
|
split_dir = os.path.join(data_dir, split)
|
||||||
os.makedirs(data_dir, exist_ok=True)
|
os.makedirs(split_dir, exist_ok=True)
|
||||||
print(f"Downloading MMLU data from {url}...")
|
path = os.path.join(split_dir, f"{subject}_{split}.csv")
|
||||||
resp = requests.get(url, stream=True, timeout=300)
|
with open(path, "w", encoding="utf-8", newline="") as f:
|
||||||
resp.raise_for_status()
|
writer = csv.writer(f)
|
||||||
total = int(resp.headers.get("content-length", 0))
|
for row in rows:
|
||||||
with tqdm.tqdm(total=total, unit="B", unit_scale=True, desc=" Download") as bar:
|
writer.writerow(row)
|
||||||
with open(tar_path, "wb") as f:
|
|
||||||
for chunk in resp.iter_content(chunk_size=8192):
|
|
||||||
f.write(chunk)
|
|
||||||
bar.update(len(chunk))
|
|
||||||
print("Extracting...")
|
|
||||||
with tarfile.open(tar_path, "r") as tf:
|
|
||||||
tf.extractall(data_dir)
|
|
||||||
os.remove(tar_path)
|
|
||||||
|
|
||||||
|
|
||||||
def download_mmlu(data_dir: str):
|
def download_mmlu(data_dir: str):
|
||||||
_download_and_extract(MMLU_URL, data_dir)
|
print(f"Downloading MMLU from HuggingFace ({MMLU_HF_DATASET}) ...")
|
||||||
src = os.path.join(data_dir, "data")
|
letters = ("A", "B", "C", "D")
|
||||||
if os.path.exists(src):
|
split_map = {"dev": "dev", "val": "validation", "test": "test"}
|
||||||
for item in os.listdir(src):
|
for local_split, hf_split in split_map.items():
|
||||||
src_item = os.path.join(src, item)
|
ds = load_dataset(MMLU_HF_DATASET, "all", split=hf_split)
|
||||||
dst_item = os.path.join(data_dir, item)
|
grouped: dict[str, list[dict]] = defaultdict(list)
|
||||||
if os.path.exists(dst_item):
|
for item in tqdm.tqdm(ds, desc=f" {local_split}", leave=False):
|
||||||
if os.path.isdir(dst_item):
|
subject = item["subject"]
|
||||||
shutil.rmtree(dst_item)
|
choices = item["choices"]
|
||||||
else:
|
ans_letter = letters[item["answer"]]
|
||||||
os.remove(dst_item)
|
grouped[subject].append(
|
||||||
os.rename(src_item, dst_item)
|
[
|
||||||
os.rmdir(src)
|
item["question"],
|
||||||
|
f"A){choices[0]}",
|
||||||
|
f"B){choices[1]}",
|
||||||
|
f"C){choices[2]}",
|
||||||
|
f"D){choices[3]}",
|
||||||
|
ans_letter,
|
||||||
|
]
|
||||||
|
)
|
||||||
|
for subject, rows in grouped.items():
|
||||||
|
_write_subject_csv(data_dir, local_split, subject, rows)
|
||||||
|
print(f" {local_split}: {len(ds)} items, {len(grouped)} subjects")
|
||||||
print(f"MMLU data saved to {data_dir}")
|
print(f"MMLU data saved to {data_dir}")
|
||||||
|
|
||||||
|
|
||||||
@@ -153,19 +155,22 @@ def build_prompt(question: str, choices: dict, subject: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def apply_chat(
|
def apply_chat(
|
||||||
tokenizer, raw_prompt: str, n_shot: int, dev_data: list[dict] | None
|
tokenizer,
|
||||||
|
raw_prompt: str,
|
||||||
|
n_shot: int,
|
||||||
|
dev_data: list[dict] | None,
|
||||||
|
subject: str = "",
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Wrap raw MMLU prompt in the model's chat template format.
|
"""Wrap raw MMLU prompt in the model's chat template format.
|
||||||
|
|
||||||
For few-shot, prepend example Q&A pairs as a second user/assistant exchange.
|
For few-shot, prepend example Q&A pairs as user/assistant exchanges.
|
||||||
|
Few-shot examples use the same subject preamble as the test question to
|
||||||
|
keep the format consistent.
|
||||||
"""
|
"""
|
||||||
messages = []
|
messages = []
|
||||||
if n_shot > 0 and dev_data:
|
if n_shot > 0 and dev_data:
|
||||||
for item in dev_data[:n_shot]:
|
for item in dev_data[:n_shot]:
|
||||||
q = f"Question: {item['question']}\n"
|
q = build_prompt(item["question"], item, subject)
|
||||||
for k in ("A", "B", "C", "D"):
|
|
||||||
q += f"{k}. {item[k]}\n"
|
|
||||||
q += "Answer:"
|
|
||||||
messages.append({"role": "user", "content": q})
|
messages.append({"role": "user", "content": q})
|
||||||
messages.append({"role": "assistant", "content": item["answer"]})
|
messages.append({"role": "assistant", "content": item["answer"]})
|
||||||
messages.append({"role": "user", "content": raw_prompt})
|
messages.append({"role": "user", "content": raw_prompt})
|
||||||
@@ -201,6 +206,25 @@ def choice_logprob(
|
|||||||
return score
|
return score
|
||||||
|
|
||||||
|
|
||||||
|
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
|
||||||
|
"""Shuffle the option order of a question.
|
||||||
|
|
||||||
|
Returns ``(permuted_item, new_answer_letter)``. The question text and
|
||||||
|
the *content* of each choice are unchanged; only which letter (A/B/C/D)
|
||||||
|
maps to which content is shuffled. This neutralises the model's
|
||||||
|
positional bias (e.g. always picking B).
|
||||||
|
"""
|
||||||
|
letters = ("A", "B", "C", "D")
|
||||||
|
contents = [item[k] for k in letters]
|
||||||
|
perm = list(letters)
|
||||||
|
rng.shuffle(perm)
|
||||||
|
permuted = {"question": item["question"]}
|
||||||
|
for new_letter, orig_letter in zip(letters, perm):
|
||||||
|
permuted[new_letter] = item[orig_letter]
|
||||||
|
new_answer = letters[perm.index(item["answer"])]
|
||||||
|
return permuted, new_answer
|
||||||
|
|
||||||
|
|
||||||
def evaluate_subject(
|
def evaluate_subject(
|
||||||
model,
|
model,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
@@ -209,18 +233,24 @@ def evaluate_subject(
|
|||||||
dev_data: list[dict] | None,
|
dev_data: list[dict] | None,
|
||||||
device: str,
|
device: str,
|
||||||
n_shot: int,
|
n_shot: int,
|
||||||
|
seed: int = 0,
|
||||||
) -> tuple[float, int, int]:
|
) -> tuple[float, int, int]:
|
||||||
|
rng = random.Random(seed) if seed >= 0 else None
|
||||||
correct = 0
|
correct = 0
|
||||||
total = 0
|
total = 0
|
||||||
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
|
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
|
||||||
raw_prompt = build_prompt(item["question"], item, subject)
|
if rng is not None:
|
||||||
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [])
|
permuted, answer = _permute_choices(item, rng)
|
||||||
|
else:
|
||||||
|
permuted, answer = item, item["answer"]
|
||||||
|
raw_prompt = build_prompt(permuted["question"], permuted, subject)
|
||||||
|
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
|
||||||
context_ids = tokenizer.encode(context)
|
context_ids = tokenizer.encode(context)
|
||||||
scores = {
|
scores = {
|
||||||
c: choice_logprob(model, tokenizer, context_ids, c, device)
|
c: choice_logprob(model, tokenizer, context_ids, c, device)
|
||||||
for c in ("A", "B", "C", "D")
|
for c in ("A", "B", "C", "D")
|
||||||
}
|
}
|
||||||
if max(scores, key=scores.get) == item["answer"]:
|
if max(scores, key=scores.get) == answer:
|
||||||
correct += 1
|
correct += 1
|
||||||
total += 1
|
total += 1
|
||||||
return correct / total, correct, total
|
return correct / total, correct, total
|
||||||
@@ -255,6 +285,12 @@ def main():
|
|||||||
default="bfloat16" if torch.cuda.is_available() else "float32",
|
default="bfloat16" if torch.cuda.is_available() else "float32",
|
||||||
help="Torch dtype",
|
help="Torch dtype",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--seed",
|
||||||
|
type=int,
|
||||||
|
default=0,
|
||||||
|
help="Seed for option permutation (0 to enable, -1 to disable)",
|
||||||
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
if args.download or not os.path.exists(args.data_dir):
|
if args.download or not os.path.exists(args.data_dir):
|
||||||
@@ -286,7 +322,14 @@ def main():
|
|||||||
test_data = load_csv(test_path)
|
test_data = load_csv(test_path)
|
||||||
|
|
||||||
acc, corr, tot = evaluate_subject(
|
acc, corr, tot = evaluate_subject(
|
||||||
model, tokenizer, subject, test_data, dev_data, device, args.n_shot
|
model,
|
||||||
|
tokenizer,
|
||||||
|
subject,
|
||||||
|
test_data,
|
||||||
|
dev_data,
|
||||||
|
device,
|
||||||
|
args.n_shot,
|
||||||
|
seed=args.seed,
|
||||||
)
|
)
|
||||||
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
||||||
total_correct += corr
|
total_correct += corr
|
||||||
|
|||||||
+96
-23
@@ -1,8 +1,10 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import json
|
import json
|
||||||
|
import time
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
from astrai.inference import InferenceEngine
|
from astrai.inference import InferenceEngine
|
||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
@@ -20,55 +22,102 @@ def processor(
|
|||||||
response_key: str,
|
response_key: str,
|
||||||
max_tokens: Optional[int],
|
max_tokens: Optional[int],
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
|
num_samples: int = 1,
|
||||||
|
cache_len: int = 2048,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
):
|
):
|
||||||
# Load model and tokenizer
|
print(f"Loading model from {param_path} ...")
|
||||||
|
t0 = time.time()
|
||||||
model = AutoModel.from_pretrained(param_path)
|
model = AutoModel.from_pretrained(param_path)
|
||||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
model.to(device="cuda", dtype=torch.bfloat16)
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
|
print(f" model loaded in {time.time() - t0:.1f}s")
|
||||||
|
|
||||||
# Create inference engine
|
|
||||||
engine = InferenceEngine(
|
engine = InferenceEngine(
|
||||||
model=model, tokenizer=tokenizer, max_batch_size=batch_size
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_batch_size=batch_size * num_samples,
|
||||||
|
max_seq_len=cache_len,
|
||||||
|
max_prompt_len=cache_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
print(f"Reading {input_json_file} ...")
|
||||||
with open(input_json_file, "r", encoding="utf-8") as f:
|
with open(input_json_file, "r", encoding="utf-8") as f:
|
||||||
input_data = [json.loads(line) for line in f]
|
input_data = [json.loads(line) for line in f]
|
||||||
|
|
||||||
# Check input format: chat messages or raw text
|
|
||||||
if input_data and "messages" in input_data[0]:
|
if input_data and "messages" in input_data[0]:
|
||||||
# Chat format: [{"messages": [...]}]
|
|
||||||
prompts = [
|
prompts = [
|
||||||
tokenizer.apply_chat_template(item["messages"], tokenize=False)
|
tokenizer.apply_chat_template(item["messages"], tokenize=False)
|
||||||
for item in input_data
|
for item in input_data
|
||||||
]
|
]
|
||||||
else:
|
else:
|
||||||
# Raw text format: [{"question": "..."}]
|
|
||||||
prompts = [item[question_key] for item in input_data]
|
prompts = [item[question_key] for item in input_data]
|
||||||
|
print(f" {len(prompts)} prompts loaded\n")
|
||||||
|
|
||||||
# Use provided max_tokens or default to model config max_len
|
|
||||||
if max_tokens is None:
|
if max_tokens is None:
|
||||||
max_tokens = model.config.max_len
|
max_tokens = model.config.max_len
|
||||||
|
|
||||||
# Generate responses (batch)
|
chunk_size = max(1, batch_size)
|
||||||
responses = engine.generate(
|
|
||||||
prompt=prompts,
|
|
||||||
stream=False,
|
|
||||||
max_tokens=max_tokens,
|
|
||||||
temperature=temperature,
|
|
||||||
top_p=top_p,
|
|
||||||
top_k=top_k,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Write results
|
|
||||||
with open(output_json_file, "w", encoding="utf-8") as f:
|
with open(output_json_file, "w", encoding="utf-8") as f:
|
||||||
for prompt, response in zip(prompts, responses):
|
pbar = tqdm(
|
||||||
if input_data and "messages" in input_data[0]:
|
total=len(prompts) * num_samples,
|
||||||
output_item = {"response": response}
|
unit="gen",
|
||||||
|
desc=f" Generating ({num_samples}x/prompt)",
|
||||||
|
)
|
||||||
|
for chunk_start in range(0, len(prompts), chunk_size):
|
||||||
|
chunk = prompts[chunk_start : chunk_start + chunk_size]
|
||||||
|
|
||||||
|
if num_samples > 1:
|
||||||
|
chunk_expanded = [p for p in chunk for _ in range(num_samples)]
|
||||||
|
resp_chunk = engine.generate(
|
||||||
|
prompt=chunk_expanded,
|
||||||
|
stream=False,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
|
)
|
||||||
|
resp_chunk = [
|
||||||
|
resp_chunk[i * num_samples : (i + 1) * num_samples]
|
||||||
|
for i in range(len(chunk))
|
||||||
|
]
|
||||||
else:
|
else:
|
||||||
output_item = {question_key: prompt, response_key: response}
|
resp_chunk = engine.generate(
|
||||||
f.write(json.dumps(output_item, ensure_ascii=False) + "\n")
|
prompt=chunk,
|
||||||
|
stream=False,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
|
)
|
||||||
|
|
||||||
|
for i, prompt in enumerate(chunk):
|
||||||
|
if input_data and "messages" in input_data[0]:
|
||||||
|
orig = input_data[chunk_start + i]
|
||||||
|
output_item = {**orig, response_key: resp_chunk[i]}
|
||||||
|
else:
|
||||||
|
output_item = {
|
||||||
|
question_key: prompt,
|
||||||
|
response_key: resp_chunk[i],
|
||||||
|
}
|
||||||
|
f.write(json.dumps(output_item, ensure_ascii=False) + "\n")
|
||||||
|
|
||||||
|
pbar.update(len(chunk) * num_samples)
|
||||||
|
|
||||||
|
pbar.close()
|
||||||
|
|
||||||
|
elapsed = time.time() - t0
|
||||||
|
print(
|
||||||
|
f"\nDone! {len(prompts)} prompts x {num_samples} samples -> {output_json_file}"
|
||||||
|
)
|
||||||
|
print(f"Total time: {elapsed:.1f}s ({elapsed / len(prompts):.2f}s/prompt)")
|
||||||
|
|
||||||
# Cleanup
|
|
||||||
engine.shutdown()
|
engine.shutdown()
|
||||||
|
|
||||||
|
|
||||||
@@ -126,12 +175,36 @@ if __name__ == "__main__":
|
|||||||
default=1,
|
default=1,
|
||||||
help="Batch size for generating responses (default: 1).",
|
help="Batch size for generating responses (default: 1).",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--num_samples",
|
||||||
|
type=int,
|
||||||
|
default=1,
|
||||||
|
help="Number of responses per prompt (expands batch internally, default: 1).",
|
||||||
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--max_tokens",
|
"--max_tokens",
|
||||||
type=int,
|
type=int,
|
||||||
default=None,
|
default=None,
|
||||||
help="Maximum tokens to generate (default: model config max_len).",
|
help="Maximum tokens to generate (default: model config max_len).",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--cache_len",
|
||||||
|
type=int,
|
||||||
|
default=2048,
|
||||||
|
help="KV cache & prompt truncation length (default: 2048, lower = less memory).",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--frequency_penalty",
|
||||||
|
type=float,
|
||||||
|
default=0.0,
|
||||||
|
help="Frequency penalty to reduce repetition (default: 0.0, try 0.5-1.0).",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--rep_window",
|
||||||
|
type=int,
|
||||||
|
default=64,
|
||||||
|
help="Window size for frequency penalty (default: 64).",
|
||||||
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
|||||||
+11
-3
@@ -8,7 +8,7 @@ import torch.optim as optim
|
|||||||
from torch import Tensor, nn
|
from torch import Tensor, nn
|
||||||
|
|
||||||
from astrai.config import AutoRegressiveLMConfig, TrainConfig
|
from astrai.config import AutoRegressiveLMConfig, TrainConfig
|
||||||
from astrai.dataset import DatasetFactory
|
from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
|
||||||
from astrai.model import AutoRegressiveLM
|
from astrai.model import AutoRegressiveLM
|
||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.trainer import SchedulerFactory, Trainer
|
from astrai.trainer import SchedulerFactory, Trainer
|
||||||
@@ -148,8 +148,8 @@ def parse_args() -> argparse.Namespace:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--max_grad_norm",
|
"--max_grad_norm",
|
||||||
type=float,
|
type=float,
|
||||||
default=1.0,
|
default=None,
|
||||||
help="Max gradient norm for clipping.",
|
help="Max gradient norm for clipping. None disables clipping.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--weight_decay",
|
"--weight_decay",
|
||||||
@@ -460,6 +460,7 @@ def train(
|
|||||||
load_path=data_root_path,
|
load_path=data_root_path,
|
||||||
window_size=window_size,
|
window_size=window_size,
|
||||||
stride=stride,
|
stride=stride,
|
||||||
|
tokenizer_path=param_path,
|
||||||
)
|
)
|
||||||
|
|
||||||
optimizer_fn = partial(
|
optimizer_fn = partial(
|
||||||
@@ -504,6 +505,12 @@ def train(
|
|||||||
|
|
||||||
grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else []
|
grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else []
|
||||||
|
|
||||||
|
collate_fn = None
|
||||||
|
if train_type == "dpo":
|
||||||
|
collate_fn = dpo_collate_fn
|
||||||
|
elif train_type == "grpo":
|
||||||
|
collate_fn = grpo_collate_fn
|
||||||
|
|
||||||
train_config = TrainConfig(
|
train_config = TrainConfig(
|
||||||
model_fn=model_fn,
|
model_fn=model_fn,
|
||||||
strategy=train_type,
|
strategy=train_type,
|
||||||
@@ -536,6 +543,7 @@ def train(
|
|||||||
executor_kwargs=executor_kwargs,
|
executor_kwargs=executor_kwargs,
|
||||||
extra_kwargs=strategy_kwargs,
|
extra_kwargs=strategy_kwargs,
|
||||||
neftune_alpha=neftune_alpha,
|
neftune_alpha=neftune_alpha,
|
||||||
|
collate_fn=collate_fn,
|
||||||
)
|
)
|
||||||
|
|
||||||
trainer = Trainer(train_config)
|
trainer = Trainer(train_config)
|
||||||
|
|||||||
+260
-48
@@ -1,14 +1,16 @@
|
|||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
|
import tempfile
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.config.preprocess_config import PipelineConfig
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
from astrai.dataset.dataset import DatasetFactory, SEQDataset
|
from astrai.dataset.dataset import DatasetFactory, dpo_tokenize
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
H5Store,
|
H5Store,
|
||||||
|
JsonlStore,
|
||||||
StoreFactory,
|
StoreFactory,
|
||||||
detect_format,
|
detect_format,
|
||||||
)
|
)
|
||||||
@@ -56,6 +58,13 @@ def _write_jsonl_dataset(test_dir, tokenizer_path, records, config_overrides=Non
|
|||||||
return data_dir
|
return data_dir
|
||||||
|
|
||||||
|
|
||||||
|
def _fake_fetch_record(self, idx, keys):
|
||||||
|
"""FakeStore.fetch_record matching real Store semantics."""
|
||||||
|
if isinstance(keys, str):
|
||||||
|
return self._data[keys][idx]
|
||||||
|
return {k: self._data[k][idx] for k in keys}
|
||||||
|
|
||||||
|
|
||||||
def _make_seq_dataset(
|
def _make_seq_dataset(
|
||||||
test_dir, name="data", seq_length=200, train_type="seq", data=None, **load_kwargs
|
test_dir, name="data", seq_length=200, train_type="seq", data=None, **load_kwargs
|
||||||
):
|
):
|
||||||
@@ -109,7 +118,7 @@ def test_dpo_strategy_with_random_data(base_test_env):
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert dpo_dataset is not None
|
assert dpo_dataset is not None
|
||||||
assert dpo_dataset.storage is not None
|
assert dpo_dataset.store is not None
|
||||||
assert len(dpo_dataset) > 0
|
assert len(dpo_dataset) > 0
|
||||||
|
|
||||||
# Test that we can get DPO items without errors
|
# Test that we can get DPO items without errors
|
||||||
@@ -138,7 +147,7 @@ def test_sft_dataset_with_random_data(base_test_env):
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert sft_dataset is not None
|
assert sft_dataset is not None
|
||||||
assert sft_dataset.storage is not None
|
assert sft_dataset.store is not None
|
||||||
assert len(sft_dataset) > 0
|
assert len(sft_dataset) > 0
|
||||||
|
|
||||||
# Test that we can get SFT items without errors
|
# Test that we can get SFT items without errors
|
||||||
@@ -169,39 +178,37 @@ def test_dataset_with_custom_stride(base_test_env):
|
|||||||
assert len(dataset) > len(default_stride_dataset)
|
assert len(dataset) > len(default_stride_dataset)
|
||||||
|
|
||||||
|
|
||||||
def test_dataset_count_property(base_test_env):
|
def test_dataset_token_count_property(base_test_env):
|
||||||
|
"""dataset.token_count exposes the raw stream token length."""
|
||||||
test_dir = base_test_env["test_dir"]
|
test_dir = base_test_env["test_dir"]
|
||||||
dataset = _make_seq_dataset(test_dir, "count_test_data")
|
dataset = _make_seq_dataset(test_dir, "count_test_data")
|
||||||
assert dataset.count == 200
|
assert dataset.token_count == 200
|
||||||
assert dataset.count > len(dataset)
|
assert dataset.token_count > len(dataset)
|
||||||
assert len(dataset) == (200 - 1 - 64) // 64 + 1
|
assert len(dataset) == (200 - 1 - 64) // 64 + 1
|
||||||
|
|
||||||
|
|
||||||
def test_empty_dataset_count():
|
|
||||||
"""Test count returns 0 when no data is loaded"""
|
|
||||||
dataset = SEQDataset(window_size=64, stride=32)
|
|
||||||
assert dataset.count == 0
|
|
||||||
assert dataset.keys == []
|
|
||||||
|
|
||||||
|
|
||||||
def test_dataset_too_short_for_window(base_test_env):
|
def test_dataset_too_short_for_window(base_test_env):
|
||||||
test_dir = base_test_env["test_dir"]
|
test_dir = base_test_env["test_dir"]
|
||||||
dataset = _make_seq_dataset(test_dir, "short", seq_length=30)
|
dataset = _make_seq_dataset(test_dir, "short", seq_length=30)
|
||||||
assert len(dataset) == 0
|
assert len(dataset) == 0
|
||||||
assert dataset.count == 30
|
assert dataset.token_count == 30
|
||||||
|
|
||||||
|
|
||||||
def test_unloaded_dataset_getitem_raises():
|
def test_unloaded_sample_window_raises():
|
||||||
"""__getitem__ without load() should fail clearly"""
|
"""Store.sample_window before load raises RuntimeError."""
|
||||||
dataset = SEQDataset(window_size=64, stride=32)
|
from astrai.dataset.storage import H5Store
|
||||||
with pytest.raises(RuntimeError, match="not loaded"):
|
|
||||||
dataset.get_index(0)
|
store = H5Store(window_size=64, stride=64)
|
||||||
|
with pytest.raises(IndexError, match="Data too short"):
|
||||||
|
store.sample_window(0)
|
||||||
|
|
||||||
|
|
||||||
def test_unloaded_dataset_len():
|
def test_unloaded_dataset_len():
|
||||||
"""__len__ without load() returns 0"""
|
"""__len__ on a store with no data returns 0."""
|
||||||
dataset = SEQDataset(window_size=64, stride=32)
|
from astrai.dataset.storage import H5Store
|
||||||
assert len(dataset) == 0
|
|
||||||
|
store = H5Store(window_size=64, stride=64)
|
||||||
|
assert len(store) == 0
|
||||||
|
|
||||||
|
|
||||||
def test_store_unloaded_len():
|
def test_store_unloaded_len():
|
||||||
@@ -214,7 +221,7 @@ def test_store_unloaded_len():
|
|||||||
def test_store_fetch_begin_equals_end(base_test_env):
|
def test_store_fetch_begin_equals_end(base_test_env):
|
||||||
test_dir = base_test_env["test_dir"]
|
test_dir = base_test_env["test_dir"]
|
||||||
dataset = _make_seq_dataset(test_dir, "empty_fetch", seq_length=100, window_size=32)
|
dataset = _make_seq_dataset(test_dir, "empty_fetch", seq_length=100, window_size=32)
|
||||||
result = dataset.storage.fetch(10, 10, "sequence")
|
result = dataset.store.fetch(10, 10, "sequence")
|
||||||
assert result.numel() == 0
|
assert result.numel() == 0
|
||||||
|
|
||||||
|
|
||||||
@@ -264,7 +271,7 @@ def test_store_multi_segment_concat(base_test_env):
|
|||||||
|
|
||||||
store = StoreFactory.create("h5")
|
store = StoreFactory.create("h5")
|
||||||
store.load(data_dir)
|
store.load(data_dir)
|
||||||
assert len(store) == 9
|
assert store.token_count == 9
|
||||||
result = store.fetch(2, 7, "sequence")
|
result = store.fetch(2, 7, "sequence")
|
||||||
assert result.tolist() == [3, 4, 5, 6, 7]
|
assert result.tolist() == [3, 4, 5, 6, 7]
|
||||||
|
|
||||||
@@ -293,7 +300,9 @@ def test_mmap_store_load_and_fetch(base_test_env):
|
|||||||
|
|
||||||
store = StoreFactory.create("bin")
|
store = StoreFactory.create("bin")
|
||||||
store.load(test_dir)
|
store.load(test_dir)
|
||||||
assert len(store) == 200
|
assert store.token_count == 200
|
||||||
|
assert store.num_records == 0
|
||||||
|
assert len(store) == 0 # no window configured, no records → 0 samples
|
||||||
assert "sequence" in store.keys
|
assert "sequence" in store.keys
|
||||||
|
|
||||||
result = store.fetch(10, 20, "sequence")
|
result = store.fetch(10, 20, "sequence")
|
||||||
@@ -306,23 +315,26 @@ def test_mmap_dataset_load(base_test_env):
|
|||||||
save_bin(test_dir, data)
|
save_bin(test_dir, data)
|
||||||
dataset = DatasetFactory.load("seq", test_dir, window_size=64)
|
dataset = DatasetFactory.load("seq", test_dir, window_size=64)
|
||||||
assert len(dataset) > 0
|
assert len(dataset) > 0
|
||||||
assert dataset.count == 200
|
assert dataset.token_count == 200
|
||||||
assert dataset[0]["input_ids"].shape[0] == 64
|
assert dataset[0]["input_ids"].shape[0] == 64
|
||||||
|
|
||||||
|
|
||||||
def test_normalize_empty_key():
|
def test_normalize_empty_key():
|
||||||
"""_normalize with empty tensor list does not crash"""
|
"""_normalize with empty tensor list does not crash."""
|
||||||
store = H5Store()
|
store = H5Store()
|
||||||
store._normalize({"sequence": []})
|
store._normalize({"sequence": []})
|
||||||
assert len(store) == 0
|
assert len(store) == 0
|
||||||
|
assert store.num_records == 0 # empty key forces num_records=0
|
||||||
assert store.keys == ["sequence"]
|
assert store.keys == ["sequence"]
|
||||||
|
|
||||||
|
|
||||||
def test_normalize_mixed_empty_key():
|
def test_normalize_mixed_empty_key():
|
||||||
"""_normalize with empty + non-empty keys returns min=0"""
|
"""_normalize with empty + non-empty keys returns min=0 records."""
|
||||||
store = H5Store()
|
store = H5Store()
|
||||||
store._normalize({"sequence": [torch.tensor([1, 2, 3])], "loss_mask": []})
|
store._normalize({"sequence": [torch.tensor([1, 2, 3])], "loss_mask": []})
|
||||||
assert len(store) == 0
|
assert len(store) == 0
|
||||||
|
assert store.num_records == 0
|
||||||
|
assert store.token_count == 0 # min() over keys
|
||||||
assert set(store.keys) == {"sequence", "loss_mask"}
|
assert set(store.keys) == {"sequence", "loss_mask"}
|
||||||
|
|
||||||
|
|
||||||
@@ -330,14 +342,14 @@ def test_grpo_dataset_dtype(base_test_env):
|
|||||||
"""GRPO dataset returns correct dtypes for per-record structured data."""
|
"""GRPO dataset returns correct dtypes for per-record structured data."""
|
||||||
from astrai.dataset.dataset import GRPODataset
|
from astrai.dataset.dataset import GRPODataset
|
||||||
|
|
||||||
test_dir = base_test_env["test_dir"]
|
|
||||||
G = 4
|
G = 4
|
||||||
dataset = GRPODataset()
|
store = type(
|
||||||
dataset.storage = type(
|
|
||||||
"FakeStore",
|
"FakeStore",
|
||||||
(),
|
(),
|
||||||
{
|
{
|
||||||
"keys": ["prompts", "responses", "masks", "rewards"],
|
"keys": ["prompts", "responses", "masks", "rewards"],
|
||||||
|
"num_records": 1,
|
||||||
|
"token_count": 0,
|
||||||
"_data": {
|
"_data": {
|
||||||
"prompts": [torch.randint(0, 100, (10,), dtype=torch.int32)],
|
"prompts": [torch.randint(0, 100, (10,), dtype=torch.int32)],
|
||||||
"responses": [
|
"responses": [
|
||||||
@@ -346,9 +358,11 @@ def test_grpo_dataset_dtype(base_test_env):
|
|||||||
"masks": [[torch.ones(5, dtype=torch.int32) for _ in range(G)]],
|
"masks": [[torch.ones(5, dtype=torch.int32) for _ in range(G)]],
|
||||||
"rewards": [torch.rand(G, dtype=torch.float32)],
|
"rewards": [torch.rand(G, dtype=torch.float32)],
|
||||||
},
|
},
|
||||||
|
"fetch_record": _fake_fetch_record,
|
||||||
|
"__len__": lambda self: self.num_records,
|
||||||
},
|
},
|
||||||
)()
|
)()
|
||||||
dataset._build_records()
|
dataset = GRPODataset(store=store)
|
||||||
item = dataset[0]
|
item = dataset[0]
|
||||||
|
|
||||||
assert item["prompts"].dtype == torch.long
|
assert item["prompts"].dtype == torch.long
|
||||||
@@ -361,25 +375,27 @@ def test_grpo_dataset_load(base_test_env):
|
|||||||
"""GRPO dataset loads record-structured data with per-response boundaries."""
|
"""GRPO dataset loads record-structured data with per-response boundaries."""
|
||||||
from astrai.dataset.dataset import GRPODataset
|
from astrai.dataset.dataset import GRPODataset
|
||||||
|
|
||||||
test_dir = base_test_env["test_dir"]
|
|
||||||
G = 3
|
G = 3
|
||||||
prompt_len = 8
|
prompt_len = 8
|
||||||
resp_lens = [5, 7, 4]
|
resp_lens = [5, 7, 4]
|
||||||
dataset = GRPODataset()
|
store = type(
|
||||||
dataset.storage = type(
|
|
||||||
"FakeStore",
|
"FakeStore",
|
||||||
(),
|
(),
|
||||||
{
|
{
|
||||||
"keys": ["prompts", "responses", "masks", "rewards"],
|
"keys": ["prompts", "responses", "masks", "rewards"],
|
||||||
|
"num_records": 1,
|
||||||
|
"token_count": 0,
|
||||||
"_data": {
|
"_data": {
|
||||||
"prompts": [torch.randint(0, 100, (prompt_len,))],
|
"prompts": [torch.randint(0, 100, (prompt_len,))],
|
||||||
"responses": [[torch.randint(0, 100, (rl,)) for rl in resp_lens]],
|
"responses": [[torch.randint(0, 100, (rl,)) for rl in resp_lens]],
|
||||||
"masks": [[torch.ones(rl, dtype=torch.int64) for rl in resp_lens]],
|
"masks": [[torch.ones(rl, dtype=torch.int64) for rl in resp_lens]],
|
||||||
"rewards": [torch.tensor([0.9, 0.3, 0.7], dtype=torch.float32)],
|
"rewards": [torch.tensor([0.9, 0.3, 0.7], dtype=torch.float32)],
|
||||||
},
|
},
|
||||||
|
"fetch_record": _fake_fetch_record,
|
||||||
|
"__len__": lambda self: self.num_records,
|
||||||
},
|
},
|
||||||
)()
|
)()
|
||||||
dataset._build_records()
|
dataset = GRPODataset(store=store)
|
||||||
|
|
||||||
assert len(dataset) == 1
|
assert len(dataset) == 1
|
||||||
item = dataset[0]
|
item = dataset[0]
|
||||||
@@ -447,16 +463,17 @@ def test_dataset_load_explicit_storage_type(base_test_env):
|
|||||||
test_dir = base_test_env["test_dir"]
|
test_dir = base_test_env["test_dir"]
|
||||||
dataset = _make_seq_dataset(test_dir, "explicit", storage_type="h5")
|
dataset = _make_seq_dataset(test_dir, "explicit", storage_type="h5")
|
||||||
assert len(dataset) > 0
|
assert len(dataset) > 0
|
||||||
assert dataset.count == 200
|
assert dataset.token_count == 200
|
||||||
|
|
||||||
|
|
||||||
def _write_json_dataset(test_dir, tokenizer_path, records, config_overrides=None):
|
def _write_json_dataset(test_dir, tokenizer_path, records, config_overrides=None):
|
||||||
"""Write JSON (not JSONL) dataset — array of objects."""
|
"""Write JSONL dataset — one JSON object per line."""
|
||||||
data_dir = os.path.join(test_dir, "json_data")
|
data_dir = os.path.join(test_dir, "json_data")
|
||||||
os.makedirs(data_dir, exist_ok=True)
|
os.makedirs(data_dir, exist_ok=True)
|
||||||
|
|
||||||
with open(os.path.join(data_dir, "data.json"), "w", encoding="utf-8") as f:
|
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
|
||||||
json.dump(records, f, ensure_ascii=False)
|
for rec in records:
|
||||||
|
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
|
||||||
|
|
||||||
config = {
|
config = {
|
||||||
"tokenizer_path": tokenizer_path,
|
"tokenizer_path": tokenizer_path,
|
||||||
@@ -535,7 +552,7 @@ def test_json_store_no_tokenizer_path(base_test_env):
|
|||||||
# Save tokenizer files directly in the dataset directory
|
# Save tokenizer files directly in the dataset directory
|
||||||
tokenizer.save_pretrained(data_dir)
|
tokenizer.save_pretrained(data_dir)
|
||||||
|
|
||||||
# Write .json data
|
# Write .jsonl data
|
||||||
records = [
|
records = [
|
||||||
{
|
{
|
||||||
"messages": [
|
"messages": [
|
||||||
@@ -544,8 +561,9 @@ def test_json_store_no_tokenizer_path(base_test_env):
|
|||||||
]
|
]
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
with open(os.path.join(data_dir, "data.json"), "w", encoding="utf-8") as f:
|
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
|
||||||
json.dump(records, f, ensure_ascii=False)
|
for rec in records:
|
||||||
|
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
|
||||||
|
|
||||||
# dataset_config.json WITHOUT tokenizer_path
|
# dataset_config.json WITHOUT tokenizer_path
|
||||||
config = {
|
config = {
|
||||||
@@ -840,11 +858,11 @@ def test_grpo_collate_variable_lengths():
|
|||||||
assert result["responses"][0, 0, 0] == 4
|
assert result["responses"][0, 0, 0] == 4
|
||||||
assert result["responses"][0, 0, 1] == 5
|
assert result["responses"][0, 0, 1] == 5
|
||||||
assert result["responses"][0, 0, 2] == 0 # padded
|
assert result["responses"][0, 0, 2] == 0 # padded
|
||||||
assert result["masks"][0, 0, 2] == False # padded
|
assert not result["masks"][0, 0, 2] # padded
|
||||||
|
|
||||||
# Check response content: item 0, response 1 is [6,7,8,9] no padding
|
# Check response content: item 0, response 1 is [6,7,8,9] no padding
|
||||||
assert result["responses"][0, 1, 3] == 9
|
assert result["responses"][0, 1, 3] == 9
|
||||||
assert result["masks"][0, 1, 3] == True
|
assert result["masks"][0, 1, 3]
|
||||||
|
|
||||||
|
|
||||||
def test_grpo_multiple_records(base_test_env):
|
def test_grpo_multiple_records(base_test_env):
|
||||||
@@ -858,12 +876,13 @@ def test_grpo_multiple_records(base_test_env):
|
|||||||
[torch.randint(0, 100, (np.random.randint(3, 8),)) for _ in range(G)]
|
[torch.randint(0, 100, (np.random.randint(3, 8),)) for _ in range(G)]
|
||||||
for _ in range(n_records)
|
for _ in range(n_records)
|
||||||
]
|
]
|
||||||
dataset = GRPODataset()
|
store = type(
|
||||||
dataset.storage = type(
|
|
||||||
"FakeStore",
|
"FakeStore",
|
||||||
(),
|
(),
|
||||||
{
|
{
|
||||||
"keys": ["prompts", "responses", "masks", "rewards"],
|
"keys": ["prompts", "responses", "masks", "rewards"],
|
||||||
|
"num_records": n_records,
|
||||||
|
"token_count": 0,
|
||||||
"_data": {
|
"_data": {
|
||||||
"prompts": [torch.randint(0, 100, (10,)) for _ in range(n_records)],
|
"prompts": [torch.randint(0, 100, (10,)) for _ in range(n_records)],
|
||||||
"responses": dummy_responses,
|
"responses": dummy_responses,
|
||||||
@@ -875,9 +894,11 @@ def test_grpo_multiple_records(base_test_env):
|
|||||||
torch.rand(G, dtype=torch.float32) for _ in range(n_records)
|
torch.rand(G, dtype=torch.float32) for _ in range(n_records)
|
||||||
],
|
],
|
||||||
},
|
},
|
||||||
|
"fetch_record": _fake_fetch_record,
|
||||||
|
"__len__": lambda self: self.num_records,
|
||||||
},
|
},
|
||||||
)()
|
)()
|
||||||
dataset._build_records()
|
dataset = GRPODataset(store=store)
|
||||||
|
|
||||||
assert len(dataset) == n_records
|
assert len(dataset) == n_records
|
||||||
|
|
||||||
@@ -888,3 +909,194 @@ def test_grpo_multiple_records(base_test_env):
|
|||||||
assert item["rewards"].shape == (G,)
|
assert item["rewards"].shape == (G,)
|
||||||
for g in range(G):
|
for g in range(G):
|
||||||
assert item["responses"][g].shape == item["masks"][g].shape
|
assert item["responses"][g].shape == item["masks"][g].shape
|
||||||
|
|
||||||
|
|
||||||
|
def _write_dpo_jsonl(test_dir, records):
|
||||||
|
"""Write a raw DPO JSONL file (no dataset_config.json)."""
|
||||||
|
path = os.path.join(test_dir, "dpo.jsonl")
|
||||||
|
with open(path, "w", encoding="utf-8") as f:
|
||||||
|
for rec in records:
|
||||||
|
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
|
||||||
|
return path
|
||||||
|
|
||||||
|
|
||||||
|
def test_dpo_tokenize_pure_function():
|
||||||
|
"""dpo_tokenize returns flat lists with correct mask alignment."""
|
||||||
|
|
||||||
|
class FakeTokenizer:
|
||||||
|
def apply_chat_template(
|
||||||
|
self, messages, tokenize=True, add_generation_prompt=True
|
||||||
|
):
|
||||||
|
ids = []
|
||||||
|
for m in messages:
|
||||||
|
ids.append(len(m["content"]))
|
||||||
|
ids.append(-1)
|
||||||
|
if add_generation_prompt:
|
||||||
|
ids.append(99)
|
||||||
|
return ids
|
||||||
|
|
||||||
|
record = {"prompt": "ab", "chosen": "xyz", "rejected": "w"}
|
||||||
|
result = dpo_tokenize(record, FakeTokenizer(), max_len=64)
|
||||||
|
|
||||||
|
assert set(result.keys()) == {"chosen", "rejected", "chosen_mask", "rejected_mask"}
|
||||||
|
assert len(result["chosen"]) == len(result["chosen_mask"])
|
||||||
|
assert len(result["rejected"]) == len(result["rejected_mask"])
|
||||||
|
|
||||||
|
assert result["chosen_mask"][0] == 0
|
||||||
|
assert any(m == 1 for m in result["chosen_mask"])
|
||||||
|
assert result["rejected_mask"][0] == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_dpo_tokenize_malformed_record():
|
||||||
|
"""dpo_tokenize returns None for missing fields."""
|
||||||
|
|
||||||
|
class FakeTokenizer:
|
||||||
|
def apply_chat_template(
|
||||||
|
self, messages, tokenize=True, add_generation_prompt=True
|
||||||
|
):
|
||||||
|
return [1]
|
||||||
|
|
||||||
|
assert dpo_tokenize({}, FakeTokenizer()) is None
|
||||||
|
assert dpo_tokenize({"prompt": "a"}, FakeTokenizer()) is None
|
||||||
|
assert dpo_tokenize({"prompt": "a", "chosen": "b"}, FakeTokenizer()) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_dpo_jsonl_lazy_load(base_test_env):
|
||||||
|
"""DPODataset loads raw JSONL with tokenizer_path → lazy processor."""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
|
||||||
|
|
||||||
|
records = [
|
||||||
|
{"input": "Hello", "chosen": "world", "rejected": "earth"},
|
||||||
|
{"input": "Foo", "chosen": "bar", "rejected": "baz"},
|
||||||
|
]
|
||||||
|
path = _write_dpo_jsonl(test_dir, records)
|
||||||
|
|
||||||
|
ds = DatasetFactory.load(
|
||||||
|
train_type="dpo",
|
||||||
|
load_path=path,
|
||||||
|
window_size=0,
|
||||||
|
tokenizer_path=tokenizer_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(ds) == 2
|
||||||
|
assert ds.store.num_records == 2
|
||||||
|
assert ds.store._processor is not None
|
||||||
|
|
||||||
|
item = ds[0]
|
||||||
|
assert set(item.keys()) == {"chosen", "rejected", "chosen_mask", "rejected_mask"}
|
||||||
|
assert item["chosen"].dtype == torch.long
|
||||||
|
assert item["chosen_mask"].dtype == torch.bool
|
||||||
|
assert item["chosen"].shape == item["chosen_mask"].shape
|
||||||
|
assert item["chosen"].shape == item["rejected"].shape
|
||||||
|
|
||||||
|
|
||||||
|
def test_dpo_jsonl_lazy_no_tokenizer():
|
||||||
|
"""DPODataset on jsonl without tokenizer_path falls back to eager
|
||||||
|
(which requires dataset_config.json, so it should raise)."""
|
||||||
|
with tempfile.TemporaryDirectory() as d:
|
||||||
|
path = os.path.join(d, "dpo.jsonl")
|
||||||
|
with open(path, "w") as f:
|
||||||
|
f.write(json.dumps({"input": "a", "chosen": "b", "rejected": "c"}) + "\n")
|
||||||
|
|
||||||
|
with pytest.raises(FileNotFoundError, match="dataset_config.json"):
|
||||||
|
DatasetFactory.load(
|
||||||
|
train_type="dpo",
|
||||||
|
load_path=path,
|
||||||
|
window_size=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_jsonl_store_lazy_len_returns_record_count(base_test_env):
|
||||||
|
"""JsonlStore in lazy mode: len() returns record count, not tokens."""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
records = [{"input": str(i), "chosen": "c", "rejected": "r"} for i in range(5)]
|
||||||
|
path = _write_dpo_jsonl(test_dir, records)
|
||||||
|
|
||||||
|
store = JsonlStore()
|
||||||
|
store.load(path, processor=lambda r: {"chosen": torch.tensor([1, 2])})
|
||||||
|
|
||||||
|
assert len(store) == 5
|
||||||
|
assert store.num_records == 5
|
||||||
|
|
||||||
|
|
||||||
|
def test_jsonl_store_eager_len_returns_token_count(base_test_env):
|
||||||
|
"""JsonlStore in eager mode: num_records reflects per-record count."""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
tokenizer_path = _save_test_tokenizer(test_dir, base_test_env["tokenizer"])
|
||||||
|
data_dir = _write_jsonl_dataset(
|
||||||
|
test_dir,
|
||||||
|
tokenizer_path,
|
||||||
|
[{"text": "hello world"}, {"text": "foo bar"}],
|
||||||
|
config_overrides={
|
||||||
|
"preprocessing": {"max_seq_len": 128, "min_chars": 0},
|
||||||
|
"output": {"position_ids_mode": "none"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
store = JsonlStore()
|
||||||
|
store.load(data_dir)
|
||||||
|
|
||||||
|
assert store.num_records == 2
|
||||||
|
assert len(store.keys) > 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_h5_store_dual_mode(base_test_env):
|
||||||
|
"""H5Store supports both fetch (stream) and fetch_record (record).
|
||||||
|
|
||||||
|
No window configured → ``len(store)`` reflects the record count
|
||||||
|
(2). ``token_count`` retains the legacy stream length (128), and
|
||||||
|
token-stream access via :meth:`fetch` is still available for
|
||||||
|
callers that want explicit begin/end control.
|
||||||
|
"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
|
||||||
|
seq_length = 64
|
||||||
|
dummy_data = {
|
||||||
|
"chosen": [_rand_seq(seq_length), _rand_seq(seq_length)],
|
||||||
|
"rejected": [_rand_seq(seq_length), _rand_seq(seq_length)],
|
||||||
|
}
|
||||||
|
save_h5(test_dir, "dpo_data", dummy_data)
|
||||||
|
|
||||||
|
store = H5Store()
|
||||||
|
store.load(test_dir)
|
||||||
|
|
||||||
|
assert store.token_count == seq_length * 2
|
||||||
|
assert store.num_records == 2
|
||||||
|
assert len(store) == 2 # no window configured → record count
|
||||||
|
|
||||||
|
rec0 = store.fetch_record(0, "chosen")
|
||||||
|
assert rec0.shape == (seq_length,)
|
||||||
|
|
||||||
|
stream = store.fetch(0, 10, "chosen")
|
||||||
|
assert stream.shape == (10,)
|
||||||
|
|
||||||
|
# Window-configured view of the same data uses stream sample count:
|
||||||
|
# token_count=128, window_size=64 → num_samples = (128-1-64)//64 + 1 = 1
|
||||||
|
stream_view = H5Store(window_size=seq_length, stride=seq_length)
|
||||||
|
stream_view.load(test_dir)
|
||||||
|
assert len(stream_view) == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_mmap_store_stream_only_no_offsets(base_test_env):
|
||||||
|
"""MmapStore without offsets: num_records == 0, stream works.
|
||||||
|
|
||||||
|
No window configured → ``len(store)`` is 0 (no iterate units).
|
||||||
|
``token_count`` remains 128 for raw token slicing, and ``fetch``
|
||||||
|
provides direct token-range access.
|
||||||
|
"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
|
||||||
|
seq_length = 128
|
||||||
|
dummy_data = {"sequence": [_rand_seq(seq_length)]}
|
||||||
|
save_bin(test_dir, dummy_data)
|
||||||
|
|
||||||
|
store = StoreFactory.create("bin")
|
||||||
|
store.load(test_dir)
|
||||||
|
|
||||||
|
assert store.token_count == seq_length
|
||||||
|
assert store.num_records == 0
|
||||||
|
assert len(store) == 0
|
||||||
|
|
||||||
|
chunk = store.fetch(0, 32, "sequence")
|
||||||
|
assert chunk.shape == (32,)
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from astrai.dataset import ResumableDistributedSampler
|
from astrai.dataset import RDSampler
|
||||||
|
|
||||||
|
|
||||||
def test_random_sampler_consistency(random_dataset):
|
def test_random_sampler_consistency(random_dataset):
|
||||||
@@ -6,8 +6,8 @@ def test_random_sampler_consistency(random_dataset):
|
|||||||
dataset = random_dataset
|
dataset = random_dataset
|
||||||
|
|
||||||
# Create two samplers with same seed
|
# Create two samplers with same seed
|
||||||
sampler1 = ResumableDistributedSampler(dataset, seed=42)
|
sampler1 = RDSampler(dataset, seed=42)
|
||||||
sampler2 = ResumableDistributedSampler(dataset, seed=42)
|
sampler2 = RDSampler(dataset, seed=42)
|
||||||
|
|
||||||
indices1 = list(iter(sampler1))
|
indices1 = list(iter(sampler1))
|
||||||
indices2 = list(iter(sampler2))
|
indices2 = list(iter(sampler2))
|
||||||
@@ -20,8 +20,8 @@ def test_random_sampler_different_seeds(random_dataset):
|
|||||||
dataset = random_dataset
|
dataset = random_dataset
|
||||||
|
|
||||||
# Create two samplers with different seeds
|
# Create two samplers with different seeds
|
||||||
sampler1 = ResumableDistributedSampler(dataset, seed=42)
|
sampler1 = RDSampler(dataset, seed=42)
|
||||||
sampler2 = ResumableDistributedSampler(dataset, seed=123)
|
sampler2 = RDSampler(dataset, seed=123)
|
||||||
|
|
||||||
indices1 = list(iter(sampler1))
|
indices1 = list(iter(sampler1))
|
||||||
indices2 = list(iter(sampler2))
|
indices2 = list(iter(sampler2))
|
||||||
@@ -35,7 +35,7 @@ def test_sampler_across_epochs(random_dataset):
|
|||||||
dataset = random_dataset
|
dataset = random_dataset
|
||||||
n = len(dataset)
|
n = len(dataset)
|
||||||
|
|
||||||
sampler = ResumableDistributedSampler(dataset, seed=42)
|
sampler = RDSampler(dataset, seed=42)
|
||||||
|
|
||||||
# Get indices for first epoch
|
# Get indices for first epoch
|
||||||
epoch1_indices = list(iter(sampler))
|
epoch1_indices = list(iter(sampler))
|
||||||
|
|||||||
Reference in New Issue
Block a user