File size: 2,063 Bytes
a94aa21
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
---
license: mit
base_model: SimpleStories/SimpleStories-V2-5M
datasets:
- desh2806/simplestories-persona-clusters-augment
tags:
- simplestories
- persona
- cluster
- supervised-fine-tuning
---

# SimpleStories cluster specialist — cluster_4

Single-epoch SFT of [SimpleStories/SimpleStories-V2-5M](https://huggingface.co/SimpleStories/SimpleStories-V2-5M) — the
**cluster_4** cluster specialist (cluster `4`) of
[desh2806/simplestories-persona-clusters-augment](https://huggingface.co/datasets/desh2806/simplestories-persona-clusters-augment), where the persona
label is the dataset's `cluster` column. One specialist is trained per cluster on
10,000 of that cluster's stories; no mixture model is trained. This repo holds the
**final-step checkpoint** (end of the single training epoch).

## Training

| | |
|---|---|
| base model | `SimpleStories/SimpleStories-V2-5M` |
| run | `cluster_4` (cluster specialist) |
| epochs | 1 (single epoch — every example seen once) |
| final step | 313 of 313 (313 steps/epoch) |
| train examples | 10000 |
| optimizer | AdamW, lr=0.0005, weight_decay=0.0 |
| batch size | 32 |
| precision | fp32 |
| seed | 42 |

Validation loss at the final checkpoint (mean cross-entropy / scored token):
- `val_own`: 1.3288
- `val_mix`: 2.2940

## Usage

```python
from transformers import AutoModelForCausalLM, AutoTokenizer

model = AutoModelForCausalLM.from_pretrained("desh2806/simplestories-cluster-cluster_4")
tokenizer = AutoTokenizer.from_pretrained("desh2806/simplestories-cluster-cluster_4")

# The base model has no BOS; seed generation with EOS (id=1) to start a new story.
import torch
seed = torch.tensor([[tokenizer.eos_token_id]])
out = model.generate(seed, max_new_tokens=150, do_sample=True, temperature=1.0, top_p=0.95,
                     eos_token_id=tokenizer.eos_token_id, pad_token_id=tokenizer.eos_token_id)
print(tokenizer.decode(out[0][1:], skip_special_tokens=True))
```

Tokenization convention used in training: `add_special_tokens=False`, EOS (id=1)
appended to every story, truncated to 512 tokens.