snip-0.4m-base / source /lora_data.py
ARotting's picture
Publish 397K parameter causal transformer pretrained from scratch
24ebd71 verified
Raw
History Blame Contribute Delete
3.22 kB
from __future__ import annotations
from itertools import product
from datasets import Dataset
from transformers import PreTrainedTokenizerFast
CONTEXT_LENGTH = 128
HEROES = [
"a careful robot",
"a brave mouse",
"a curious child",
"a small fox",
"a patient inventor",
"a lonely star",
]
PLACES = [
"a moonlit castle",
"a clockwork garden",
"a floating library",
"a quiet workshop",
"a crystal forest",
"an underwater city",
]
GOALS = [
"find a lost key",
"repair a broken bridge",
"help a frightened friend",
"learn why the bells stopped",
"return a borrowed light",
]
LESSONS = [
"courage can be quiet",
"asking for help is wise",
"patience can solve hard problems",
"kindness changes a whole journey",
"mistakes can become maps",
]
def build_examples() -> list[dict[str, str]]:
examples: list[dict[str, str]] = []
for index, (hero, place, goal, lesson) in enumerate(
product(HEROES, PLACES, GOALS, LESSONS)
):
object_name = ["lantern", "silver thread", "paper crown", "tiny compass"][index % 4]
prompt = (
f"Write a tiny story about {hero} in {place}. "
f"The hero must {goal} and learn that {lesson}."
)
response = (
f"In {place}, {hero} carried a {object_name}. The path seemed impossible, "
f"but the hero chose to {goal}. A new friend noticed the effort and offered "
f"one small clue. Together they finished before sunrise. From then on, the "
f"hero remembered that {lesson}."
)
examples.append({"prompt": prompt, "response": response})
return examples
def split_examples() -> tuple[list[dict[str, str]], list[dict[str, str]]]:
examples = build_examples()
train = [example for index, example in enumerate(examples) if index % 10 != 0]
evaluation = [example for index, example in enumerate(examples) if index % 10 == 0]
return train, evaluation
def encode_examples(
examples: list[dict[str, str]],
tokenizer: PreTrainedTokenizerFast,
) -> Dataset:
rows = {"input_ids": [], "attention_mask": [], "labels": []}
for example in examples:
prefix = f"<bos>User: {example['prompt']}\nAssistant:"
full_text = f"{prefix} {example['response']}<eos>"
full = tokenizer(
full_text,
max_length=CONTEXT_LENGTH,
truncation=True,
padding="max_length",
add_special_tokens=False,
)
prefix_ids = tokenizer(
prefix,
max_length=CONTEXT_LENGTH,
truncation=True,
add_special_tokens=False,
)["input_ids"]
labels = list(full["input_ids"])
masked_prefix = min(len(prefix_ids), CONTEXT_LENGTH)
labels[:masked_prefix] = [-100] * masked_prefix
labels = [
label if attention else -100
for label, attention in zip(labels, full["attention_mask"], strict=True)
]
rows["input_ids"].append(full["input_ids"])
rows["attention_mask"].append(full["attention_mask"])
rows["labels"].append(labels)
return Dataset.from_dict(rows)