FLustc's picture
Upload folder using huggingface_hub
3194a67 verified
Raw
History Blame Contribute Delete
2.31 kB
from __future__ import annotations
from pathlib import Path
import torch
from torch.utils.data import Dataset
from .tokenizer import CharacterTokenizer
FALLBACK_LINES = [
"the quick brown fox jumps over the lazy dog",
"machine learning models can generate text",
"diffusion models gradually remove noise",
"multi mask diffusion preserves more information",
"language models predict clean tokens from noisy inputs",
"deep learning requires data optimization and evaluation",
"multiple masks retain a trace of the original token",
"a small reproducible experiment is better than no result",
"time conditioned transformers learn iterative denoising",
"research code should save reload and evaluate checkpoints",
"clean tokens are assigned to designated mask classes",
"masked language modeling reconstructs hidden symbols",
"simple datasets make the training pipeline dependable",
"few step generation trades computation for model quality",
"the model observes noisy sequences and predicts characters",
"careful tests catch errors before expensive gpu training",
]
def load_corpus(path: str | Path, minimum_lines: int = 512) -> list[str]:
path = Path(path)
if path.exists():
lines = [
line.strip().lower()
for line in path.read_text(encoding="utf-8").splitlines()
if line.strip()
]
else:
lines = list(FALLBACK_LINES)
if not lines:
lines = list(FALLBACK_LINES)
repeats = (minimum_lines + len(lines) - 1) // len(lines)
expanded = (lines * repeats)[:minimum_lines]
return [f"{line}\n" for line in expanded]
class TextDataset(Dataset):
def __init__(
self,
texts: list[str],
tokenizer: CharacterTokenizer,
seq_len: int,
):
self.texts = texts
self.tokenizer = tokenizer
self.seq_len = seq_len
def __len__(self) -> int:
return len(self.texts)
def __getitem__(self, index: int) -> dict[str, torch.Tensor]:
ids = torch.tensor(
self.tokenizer.encode(self.texts[index], self.seq_len),
dtype=torch.long,
)
attention_mask = ids.ne(self.tokenizer.pad_id)
return {"input_ids": ids, "attention_mask": attention_mask}