Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
@@ -0,0 +1,355 @@
|
||||
"""Pretrain / SFT sample construction.
|
||||
|
||||
Pretrain: Wikipedia parquet → tokenize → pack (B, T). Languages mix 1:1 by
|
||||
token via seq_len-sized blocks so each training chunk is monolingual.
|
||||
|
||||
SFT: instruction-parallel rows → prompt-masked labels. Template lives in
|
||||
``prompts.instruction_prompt`` (same string as eval_mt).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Iterable, Protocol
|
||||
|
||||
import torch
|
||||
|
||||
from .prompts import instruction_prompt
|
||||
|
||||
WIKI_SHARD_TOTAL = {"zh": 6, "en": 41}
|
||||
WIKI_BASE = (
|
||||
"https://huggingface.co/datasets/wikimedia/wikipedia/resolve/main/20231101.{lang}"
|
||||
)
|
||||
IGNORE_INDEX = -100
|
||||
|
||||
|
||||
class Tokenizer(Protocol):
|
||||
vocab_size: int
|
||||
|
||||
def encode(self, text: str) -> list[int]: ...
|
||||
|
||||
def decode(self, ids: list[int]) -> str: ...
|
||||
|
||||
|
||||
@dataclass
|
||||
class SentencePieceTokenizer:
|
||||
_sp: object
|
||||
|
||||
@property
|
||||
def vocab_size(self) -> int:
|
||||
return int(self._sp.vocab_size())
|
||||
|
||||
def encode(self, text: str) -> list[int]:
|
||||
return list(self._sp.encode(text, out_type=int))
|
||||
|
||||
def decode(self, ids: list[int]) -> str:
|
||||
return str(self._sp.decode(ids))
|
||||
|
||||
|
||||
@dataclass
|
||||
class HuggingFaceTokenizer:
|
||||
_tok: object
|
||||
|
||||
@property
|
||||
def vocab_size(self) -> int:
|
||||
return int(len(self._tok))
|
||||
|
||||
def encode(self, text: str) -> list[int]:
|
||||
return list(self._tok.encode(text, add_special_tokens=False))
|
||||
|
||||
def decode(self, ids: list[int]) -> str:
|
||||
return str(self._tok.decode(ids, skip_special_tokens=True))
|
||||
|
||||
|
||||
def load_tokenizer(source: str) -> Tokenizer:
|
||||
"""`.model` 走 SentencePiece, 其它当作 HuggingFace 名或本地目录."""
|
||||
if source.endswith(".model"):
|
||||
from sentencepiece import SentencePieceProcessor
|
||||
|
||||
return SentencePieceTokenizer(SentencePieceProcessor(model_file=source))
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
tok = AutoTokenizer.from_pretrained(source, trust_remote_code=True)
|
||||
return HuggingFaceTokenizer(tok)
|
||||
|
||||
|
||||
def pretrain_dir() -> Path:
|
||||
for candidate in (
|
||||
os.environ.get("KDA_PRETRAIN_DIR"),
|
||||
"/data/pretrain",
|
||||
"data/pretrain",
|
||||
):
|
||||
if candidate and Path(candidate).is_dir():
|
||||
return Path(candidate)
|
||||
return Path("data/pretrain")
|
||||
|
||||
|
||||
def _wiki_files(lang: str, n_shards: int) -> list[str]:
|
||||
if lang not in WIKI_SHARD_TOTAL:
|
||||
raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en")
|
||||
total = WIKI_SHARD_TOTAL[lang]
|
||||
n = min(max(n_shards, 1), total)
|
||||
base = WIKI_BASE.format(lang=lang)
|
||||
return [f"{base}/train-{i:05d}-of-{total:05d}.parquet" for i in range(n)]
|
||||
|
||||
|
||||
def _cache_path(cache_dir: Path, lang: str, n_shards: int, limit: int) -> Path:
|
||||
return cache_dir / f"wiki-{lang}-n{n_shards}-limit{limit}.jsonl"
|
||||
|
||||
|
||||
def fetch_wiki_texts(
|
||||
limit: int,
|
||||
lang: str = "zh",
|
||||
n_shards: int = 2,
|
||||
cache_dir: str | Path | None = None,
|
||||
) -> list[str]:
|
||||
"""Load up to ``limit`` article bodies, caching jsonl under pretrain_dir."""
|
||||
cache = Path(cache_dir) if cache_dir is not None else pretrain_dir()
|
||||
cache.mkdir(parents=True, exist_ok=True)
|
||||
path = _cache_path(cache, lang, n_shards, limit)
|
||||
if path.exists():
|
||||
texts: list[str] = []
|
||||
with path.open(encoding="utf-8") as fh:
|
||||
for line in fh:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
texts.append(json.loads(line)["text"])
|
||||
if len(texts) >= limit:
|
||||
break
|
||||
if texts:
|
||||
return texts
|
||||
|
||||
from datasets import load_dataset
|
||||
|
||||
files = _wiki_files(lang, n_shards)
|
||||
ds = load_dataset("parquet", data_files=files, split="train", streaming=True)
|
||||
texts = []
|
||||
for i, row in enumerate(ds):
|
||||
if i >= limit:
|
||||
break
|
||||
texts.append(row["text"])
|
||||
tmp = path.with_suffix(path.suffix + ".tmp")
|
||||
with tmp.open("w", encoding="utf-8") as fh:
|
||||
for text in texts:
|
||||
fh.write(json.dumps({"text": text}, ensure_ascii=False) + "\n")
|
||||
tmp.replace(path)
|
||||
return texts
|
||||
|
||||
|
||||
def tokenize_corpus(texts: list[str], tok: Tokenizer) -> list[int]:
|
||||
ids: list[int] = []
|
||||
for text in texts:
|
||||
ids.extend(tok.encode(text))
|
||||
return ids
|
||||
|
||||
|
||||
def interleave_balanced(ids_a: list[int], ids_b: list[int], block: int) -> list[int]:
|
||||
"""1:1 by token: seq_len-sized monolingual blocks, drop the longer tail."""
|
||||
if block < 1:
|
||||
raise ValueError(f"block must be >= 1, got {block}")
|
||||
n = min(len(ids_a), len(ids_b))
|
||||
n = (n // block) * block
|
||||
out: list[int] = []
|
||||
a, b = ids_a, ids_b
|
||||
for i in range(0, n, block):
|
||||
out.extend(a[i : i + block])
|
||||
out.extend(b[i : i + block])
|
||||
return out
|
||||
|
||||
|
||||
def chunk_ids(ids: list[int], batch: int, seq_len: int) -> torch.Tensor:
|
||||
"""切成 (num_chunks, B, T); 末尾不足部分丢弃."""
|
||||
n = (len(ids) // (batch * seq_len)) * (batch * seq_len)
|
||||
t = torch.tensor(ids[:n], dtype=torch.long)
|
||||
if n == 0:
|
||||
return t.view(0, batch, seq_len)
|
||||
return t.view(batch, -1, seq_len).transpose(0, 1)
|
||||
|
||||
|
||||
def split_heldout(
|
||||
chunks: torch.Tensor,
|
||||
frac: float = 0.01,
|
||||
min_heldout: int = 1,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Last ``frac`` of packed chunks for CE only. Empty held-out if too few."""
|
||||
n = int(chunks.size(0))
|
||||
if n <= 1 or frac <= 0:
|
||||
return chunks, chunks[:0]
|
||||
h = max(min_heldout, int(n * frac))
|
||||
h = min(h, n - 1)
|
||||
return chunks[:-h], chunks[-h:]
|
||||
|
||||
|
||||
def iter_chunks(chunks: torch.Tensor):
|
||||
"""逐块产出 (input_ids, labels), labels 右移 (模型内 CE shift)."""
|
||||
for chunk in chunks:
|
||||
yield chunk, chunk.clone()
|
||||
|
||||
|
||||
def iter_indexed(chunks: torch.Tensor, start: int = 0):
|
||||
"""Infinite cycle with a global index (for --resume)."""
|
||||
n = int(chunks.size(0))
|
||||
if n == 0:
|
||||
raise ValueError("no training chunks")
|
||||
i = start
|
||||
while True:
|
||||
x = chunks[i % n]
|
||||
yield i, x, x.clone()
|
||||
i += 1
|
||||
|
||||
|
||||
def load_pretrain_chunks(
|
||||
tok: Tokenizer,
|
||||
*,
|
||||
langs: Iterable[str],
|
||||
limit: int,
|
||||
batch: int,
|
||||
seq_len: int,
|
||||
heldout_frac: float = 0.01,
|
||||
n_shards: int = 2,
|
||||
cache_dir: str | Path | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, int]:
|
||||
"""Fetch / cache / tokenize / pack. Returns train chunks, held-out, token count."""
|
||||
lang_list = [lang.strip() for lang in langs if lang.strip()]
|
||||
if not lang_list:
|
||||
raise ValueError("langs must contain at least one of zh, en")
|
||||
streams: list[list[int]] = []
|
||||
for lang in lang_list:
|
||||
print(f"loading {limit} wiki articles ({lang}) ...")
|
||||
texts = fetch_wiki_texts(limit, lang=lang, n_shards=n_shards, cache_dir=cache_dir)
|
||||
streams.append(tokenize_corpus(texts, tok))
|
||||
print(f" {lang}: {len(streams[-1]):,} tokens from {len(texts)} articles")
|
||||
if len(streams) == 1:
|
||||
ids = streams[0]
|
||||
else:
|
||||
ids = streams[0]
|
||||
for extra in streams[1:]:
|
||||
ids = interleave_balanced(ids, extra, seq_len)
|
||||
chunks = chunk_ids(ids, batch, seq_len)
|
||||
train, held = split_heldout(chunks, heldout_frac)
|
||||
return train, held, len(ids)
|
||||
|
||||
|
||||
def pad_id(tok: Tokenizer) -> int:
|
||||
inner = getattr(tok, "_tok", None)
|
||||
if inner is not None:
|
||||
pid = getattr(inner, "pad_token_id", None)
|
||||
if pid is not None:
|
||||
return int(pid)
|
||||
eid = getattr(inner, "eos_token_id", None)
|
||||
if eid is not None:
|
||||
return int(eid)
|
||||
return 0
|
||||
|
||||
|
||||
def eos_id(tok: Tokenizer) -> int | None:
|
||||
inner = getattr(tok, "_tok", None)
|
||||
if inner is not None:
|
||||
eid = getattr(inner, "eos_token_id", None)
|
||||
if eid is not None:
|
||||
return int(eid)
|
||||
convert = getattr(inner, "convert_tokens_to_ids", None)
|
||||
if convert is not None:
|
||||
tid = convert("<|im_end|>")
|
||||
if isinstance(tid, int) and tid >= 0:
|
||||
return tid
|
||||
return None
|
||||
|
||||
|
||||
def encode_sft_row(
|
||||
tok: Tokenizer,
|
||||
src: str,
|
||||
tgt: str,
|
||||
target_lang: str,
|
||||
max_len: int,
|
||||
eos: int | None = None,
|
||||
) -> tuple[list[int], list[int]]:
|
||||
prompt_ids = tok.encode(instruction_prompt(src, target_lang))
|
||||
tgt_ids = tok.encode(tgt)
|
||||
if eos is not None:
|
||||
tgt_ids = tgt_ids + [eos]
|
||||
ids = prompt_ids + tgt_ids
|
||||
labels = [IGNORE_INDEX] * len(prompt_ids) + list(tgt_ids)
|
||||
if len(ids) > max_len:
|
||||
overflow = len(ids) - max_len
|
||||
cut = min(overflow, max(len(prompt_ids) - 1, 0))
|
||||
ids = ids[cut:]
|
||||
labels = labels[cut:]
|
||||
if len(ids) > max_len:
|
||||
ids = ids[:max_len]
|
||||
labels = labels[:max_len]
|
||||
return ids, labels
|
||||
|
||||
|
||||
def load_sft_rows(path: str | Path) -> list[dict]:
|
||||
"""jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``."""
|
||||
p = Path(path)
|
||||
rows: list[dict] = []
|
||||
text = p.read_text(encoding="utf-8")
|
||||
if p.suffix == ".jsonl" or p.suffix == ".json":
|
||||
for line in text.splitlines():
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
obj = json.loads(line)
|
||||
rows.append(
|
||||
{
|
||||
"src": obj["src"],
|
||||
"tgt": obj["tgt"],
|
||||
"target_lang": obj.get("target_lang", "en"),
|
||||
}
|
||||
)
|
||||
return rows
|
||||
for line in text.splitlines():
|
||||
line = line.strip()
|
||||
if not line or line.startswith("#"):
|
||||
continue
|
||||
parts = line.split("\t")
|
||||
if len(parts) < 2:
|
||||
raise ValueError(f"SFT TSV needs src, tgt [, target_lang]: {line[:80]!r}")
|
||||
lang = parts[2] if len(parts) > 2 else "en"
|
||||
rows.append({"src": parts[0], "tgt": parts[1], "target_lang": lang})
|
||||
return rows
|
||||
|
||||
|
||||
def collate_sft(
|
||||
rows: list[dict],
|
||||
tok: Tokenizer,
|
||||
max_len: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
pad = pad_id(tok)
|
||||
eos = eos_id(tok)
|
||||
encoded = [
|
||||
encode_sft_row(tok, r["src"], r["tgt"], r["target_lang"], max_len, eos)
|
||||
for r in rows
|
||||
]
|
||||
width = min(max(len(ids) for ids, _ in encoded), max_len)
|
||||
width = max(width, 2)
|
||||
bsz = len(encoded)
|
||||
input_ids = torch.full((bsz, width), pad, dtype=torch.long)
|
||||
labels = torch.full((bsz, width), IGNORE_INDEX, dtype=torch.long)
|
||||
for i, (ids, lab) in enumerate(encoded):
|
||||
n = min(len(ids), width)
|
||||
input_ids[i, :n] = torch.tensor(ids[:n], dtype=torch.long)
|
||||
labels[i, :n] = torch.tensor(lab[:n], dtype=torch.long)
|
||||
return input_ids, labels
|
||||
|
||||
|
||||
def iter_sft_batches(
|
||||
rows: list[dict],
|
||||
tok: Tokenizer,
|
||||
batch: int,
|
||||
max_len: int,
|
||||
start: int = 0,
|
||||
):
|
||||
n = len(rows)
|
||||
if n == 0:
|
||||
raise ValueError("no SFT rows")
|
||||
i = start
|
||||
while True:
|
||||
sl = [rows[j % n] for j in range(i, i + batch)]
|
||||
yield i, *collate_sft(sl, tok, max_len)
|
||||
i += batch
|
||||
Reference in New Issue
Block a user