| import json |
| import re |
| from dataclasses import dataclass |
| from pathlib import Path |
|
|
| import torch |
| from tokenizers import Tokenizer |
|
|
| from model import GPT, GPTConfig |
|
|
|
|
| EOS_TOKEN = "<|endoftext|>" |
| DEFAULT_PROMPT = "Once upon a time" |
| DEFAULT_TARGET_TOKENS = 120 |
| DEFAULT_EXTRA_TOKENS = 80 |
| DEFAULT_TEMPERATURE = 0.8 |
| DEFAULT_TOP_K = 40 |
|
|
| MAX_PROMPT_TOKENS = 256 |
| MAX_TARGET_TOKENS = 500 |
| MAX_EXTRA_TOKENS = 200 |
| MIN_TEMPERATURE = 0.1 |
| MAX_TEMPERATURE = 2.0 |
|
|
| END_PUNCTUATION = (".", "!", "?") |
| STORY_START_PATTERN = re.compile( |
| r"\b(?:once upon a time|there was once|there once was)\b", |
| re.IGNORECASE, |
| ) |
|
|
|
|
| @dataclass(frozen=True) |
| class GenerationResult: |
| story: str |
| generated_tokens: int |
|
|
|
|
| def get_device() -> torch.device: |
| if torch.backends.mps.is_available(): |
| return torch.device("mps") |
|
|
| if torch.cuda.is_available(): |
| return torch.device("cuda") |
|
|
| return torch.device("cpu") |
|
|
|
|
| def _build_config(config_data: object) -> GPTConfig: |
| if isinstance(config_data, GPTConfig): |
| return config_data |
|
|
| if isinstance(config_data, dict): |
| return GPTConfig(**config_data) |
|
|
| if hasattr(config_data, "__dict__"): |
| return GPTConfig(**vars(config_data)) |
|
|
| raise ValueError("Checkpoint contains an unsupported model configuration.") |
|
|
|
|
| def load_training_checkpoint( |
| checkpoint_path: str | Path, |
| device: torch.device, |
| ) -> GPT: |
| checkpoint = torch.load( |
| Path(checkpoint_path), |
| map_location="cpu", |
| weights_only=False, |
| ) |
|
|
| if not isinstance(checkpoint, dict): |
| raise ValueError("Checkpoint must contain a dictionary.") |
|
|
| if "model_state" not in checkpoint or "config" not in checkpoint: |
| raise ValueError("Checkpoint is missing model_state or config.") |
|
|
| model = GPT(_build_config(checkpoint["config"])) |
| model.load_state_dict(checkpoint["model_state"]) |
| model.to(device) |
| model.eval() |
| return model |
|
|
|
|
| def load_exported_model( |
| config_path: str | Path, |
| weights_path: str | Path, |
| device: torch.device, |
| ) -> GPT: |
| config_data = json.loads(Path(config_path).read_text(encoding="utf-8")) |
| model = GPT(_build_config(config_data)) |
| state_dict = torch.load( |
| Path(weights_path), |
| map_location="cpu", |
| weights_only=True, |
| ) |
|
|
| if not isinstance(state_dict, dict): |
| raise ValueError("Exported weights must contain a state dictionary.") |
|
|
| model.load_state_dict(state_dict) |
| model.to(device) |
| model.eval() |
| return model |
|
|
|
|
| def normalize_text(text: str) -> str: |
| text = text.replace(EOS_TOKEN, "") |
| text = re.sub(r"\s+", " ", text) |
| text = re.sub(r"\s+([,.;:!?])", r"\1", text) |
| return text.strip() |
|
|
|
|
| def ends_with_sentence(text: str) -> bool: |
| text = normalize_text(text) |
| return bool(re.search(r"""[.!?](?:["'\u2019\u201d])?$""", text)) |
|
|
|
|
| def trim_repeated_story(text: str, prompt: str = "") -> str: |
| text = normalize_text(text) |
| normalized_prompt = normalize_text(prompt) |
| prompt_boundary = ( |
| len(normalized_prompt) if text.startswith(normalized_prompt) else 0 |
| ) |
|
|
| for match in STORY_START_PATTERN.finditer(text): |
| if match.start() < prompt_boundary or match.start() == 0: |
| continue |
|
|
| candidate = text[: match.start()].strip() |
| if len(candidate.split()) >= 20: |
| return candidate |
|
|
| return text |
|
|
|
|
| def trim_to_last_sentence(text: str, prompt: str = "") -> str: |
| text = normalize_text(text) |
| normalized_prompt = normalize_text(prompt) |
| last_position = max(text.rfind(mark) for mark in END_PUNCTUATION) |
|
|
| if last_position == -1: |
| return text |
|
|
| if text.startswith(normalized_prompt) and last_position < len(normalized_prompt): |
| return text |
|
|
| return text[: last_position + 1].strip() |
|
|
|
|
| def clean_story(text: str, prompt: str = "") -> str: |
| text = trim_repeated_story(text, prompt=prompt) |
| text = trim_to_last_sentence(text, prompt=prompt) |
| return normalize_text(text) |
|
|
|
|
| def validate_generation_inputs( |
| tokenizer: Tokenizer, |
| prompt: object, |
| target_tokens: object, |
| extra_tokens: object, |
| temperature: object, |
| top_k: object, |
| ) -> tuple[str, list[int], int, int, float, int]: |
| if not isinstance(prompt, str): |
| raise ValueError("prompt must be a string.") |
|
|
| if type(target_tokens) is not int: |
| raise ValueError("tokens must be an integer.") |
|
|
| if type(extra_tokens) is not int: |
| raise ValueError("extra_tokens must be an integer.") |
|
|
| if isinstance(temperature, bool) or not isinstance(temperature, (int, float)): |
| raise ValueError("temperature must be a number.") |
|
|
| if type(top_k) is not int: |
| raise ValueError("top_k must be an integer.") |
|
|
| if not 1 <= target_tokens <= MAX_TARGET_TOKENS: |
| raise ValueError(f"tokens must be between 1 and {MAX_TARGET_TOKENS}.") |
|
|
| if not 0 <= extra_tokens <= MAX_EXTRA_TOKENS: |
| raise ValueError( |
| f"extra_tokens must be between 0 and {MAX_EXTRA_TOKENS}." |
| ) |
|
|
| temperature = float(temperature) |
| if not MIN_TEMPERATURE <= temperature <= MAX_TEMPERATURE: |
| raise ValueError( |
| f"temperature must be between {MIN_TEMPERATURE} and {MAX_TEMPERATURE}." |
| ) |
|
|
| vocab_size = tokenizer.get_vocab_size() |
| if not 1 <= top_k <= vocab_size: |
| raise ValueError(f"top_k must be between 1 and {vocab_size}.") |
|
|
| prompt = prompt.strip() or DEFAULT_PROMPT |
| prompt_ids = tokenizer.encode(prompt).ids |
|
|
| if not prompt_ids: |
| raise ValueError("prompt must contain text.") |
|
|
| if len(prompt_ids) > MAX_PROMPT_TOKENS: |
| raise ValueError( |
| f"prompt must not exceed {MAX_PROMPT_TOKENS} encoded tokens." |
| ) |
|
|
| return ( |
| prompt, |
| prompt_ids, |
| target_tokens, |
| extra_tokens, |
| temperature, |
| top_k, |
| ) |
|
|
|
|
| def sample_next_token( |
| model: GPT, |
| input_ids: torch.Tensor, |
| temperature: float, |
| top_k: int, |
| ) -> torch.Tensor: |
| idx_cond = input_ids[:, -model.config.block_size :] |
| logits, _ = model(idx_cond) |
| logits = logits[:, -1, :] / temperature |
|
|
| values, _ = torch.topk(logits, min(top_k, logits.size(-1))) |
| logits = logits.masked_fill(logits < values[:, [-1]], float("-inf")) |
|
|
| probabilities = torch.softmax(logits, dim=-1) |
| return torch.multinomial(probabilities, num_samples=1) |
|
|
|
|
| @torch.no_grad() |
| def generate_story( |
| model: GPT, |
| tokenizer: Tokenizer, |
| prompt: object = DEFAULT_PROMPT, |
| target_tokens: object = DEFAULT_TARGET_TOKENS, |
| extra_tokens: object = DEFAULT_EXTRA_TOKENS, |
| temperature: object = DEFAULT_TEMPERATURE, |
| top_k: object = DEFAULT_TOP_K, |
| device: torch.device | None = None, |
| ) -> GenerationResult: |
| ( |
| prompt, |
| prompt_ids, |
| target_tokens, |
| extra_tokens, |
| temperature, |
| top_k, |
| ) = validate_generation_inputs( |
| tokenizer=tokenizer, |
| prompt=prompt, |
| target_tokens=target_tokens, |
| extra_tokens=extra_tokens, |
| temperature=temperature, |
| top_k=top_k, |
| ) |
|
|
| if device is None: |
| device = next(model.parameters()).device |
|
|
| input_ids = torch.tensor([prompt_ids], dtype=torch.long, device=device) |
| eos_token_id = tokenizer.token_to_id(EOS_TOKEN) |
| generated_tokens = 0 |
|
|
| for generated_tokens in range(1, target_tokens + extra_tokens + 1): |
| next_id = sample_next_token( |
| model=model, |
| input_ids=input_ids, |
| temperature=temperature, |
| top_k=top_k, |
| ) |
| input_ids = torch.cat((input_ids, next_id), dim=1) |
|
|
| if eos_token_id is not None and next_id.item() == eos_token_id: |
| break |
|
|
| if generated_tokens >= target_tokens: |
| current_text = tokenizer.decode( |
| input_ids[0].tolist(), |
| skip_special_tokens=False, |
| ) |
| if ends_with_sentence(current_text): |
| break |
|
|
| generated_ids = input_ids[0].tolist() |
| if eos_token_id is not None and eos_token_id in generated_ids: |
| generated_ids = generated_ids[: generated_ids.index(eos_token_id)] |
|
|
| text = tokenizer.decode(generated_ids, skip_special_tokens=False) |
| story = clean_story(text, prompt=prompt) |
|
|
| if not story: |
| raise RuntimeError("The model generated an empty result.") |
|
|
| return GenerationResult( |
| story=story, |
| generated_tokens=generated_tokens, |
| ) |
|
|