File size: 2,167 Bytes
2b2f0da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
50f1836
2b2f0da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
50f1836
 
 
 
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-personas-10k
tags:
- simplestories
- persona
- supervised-fine-tuning
---

# SimpleStories persona model — mixture

Single-epoch SFT of [SimpleStories/SimpleStories-V2-5M](https://huggingface.co/SimpleStories/SimpleStories-V2-5M)
on the **uniform mixture (union of all 5 personas)** from [desh2806/simplestories-personas-10k](https://huggingface.co/datasets/desh2806/simplestories-personas-10k).

Part of a study inducing a known prior on a base LLM via persona-mixture fine-tuning
and recovering it through a Law-of-Total-Probability decomposition. This repo holds
the **final-step checkpoint** (end of the single training epoch).

## Training

| | |
|---|---|
| base model | `SimpleStories/SimpleStories-V2-5M` |
| run | `mixture` (uniform mixture (union of all 5 personas)) |
| epochs | 1 (single epoch — every example seen once) |
| final step | 1421 of 1421 (1421 steps/epoch) |
| train examples | 45472 |
| 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_mix`: 1.7983

## Usage

```python
from transformers import AutoModelForCausalLM, AutoTokenizer

model = AutoModelForCausalLM.from_pretrained("desh2806/simplestories-persona-mixture")
tokenizer = AutoTokenizer.from_pretrained("desh2806/simplestories-persona-mixture")

# 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`, every story
wrapped in EOS (id=1) on both sides — `[EOS, tokens…, EOS]` — truncated to 512
tokens. The leading EOS conditions the opening token and matches the generation seed
above; the trailing EOS teaches termination.