Spaces:
Sleeping
Sleeping
| """Stable prompt and artifact contract shared by the Studio and product SFT.""" | |
| from __future__ import annotations | |
| import hashlib | |
| import json | |
| import re | |
| from pathlib import Path | |
| from typing import Any, Mapping, Sequence | |
| CHAT_PROMPT_FORMAT = "smolgpt-fables-chat-v3" | |
| SMOLLM3_CHAT_PROMPT_FORMAT = "smolgpt-fables-smollm3-chat-v5" | |
| RAW_PROMPT_FORMAT = "smolgpt-fables-raw-v1" | |
| PRODUCT_CONTEXT_LENGTH = 2048 | |
| SMOLLM3_PRODUCT_CONTEXT_LENGTH = 4096 | |
| CHAT_SYSTEM_PROMPT = ( | |
| "You are SmolGPT-Fables. Output only the requested finished Markdown story " | |
| "continuation. Begin exactly with `### Scene 01:`. Emit exactly the requested " | |
| "number of consecutive, zero-padded `### Scene NN:` sections. Those scene " | |
| "headings are the only headings allowed in the output: never emit an H1, H2, " | |
| "any other H3, an H4 or deeper heading, or any other section before, between, " | |
| "or after them. Follow the canvas exactly and copy every required name and " | |
| "detail verbatim. Stop immediately after the requested final scene. Never " | |
| "repeat or quote the story title, metadata, canvas, `## Story`, or any instruction." | |
| ) | |
| CHAT_PROMPT_TRANSFORM_VERSION = "scene-output-contract-v1" | |
| CHAT_STORY_BOUNDARY = "## Story\n\n" | |
| CHAT_OUTPUT_CONTRACT_TEMPLATE = ( | |
| "Output contract (follow exactly): use only these H3 heading prefixes, in this " | |
| "order: {headings}. Write exactly {scene_count} scenes; emit no other heading " | |
| "or section; stop immediately after completing `### Scene {final_scene:02d}:`.\n\n" | |
| ) | |
| SMOLLM3_PROMPT_TRANSFORM_VERSION = "natural-fable-scene-contract-v3" | |
| SMOLLM3_STORY_BOUNDARY = "## Story\n\n" | |
| SMOLLM3_SCENE_WORD_RANGE = (45, 115) | |
| SMOLLM3_SYSTEM_PROMPT = ( | |
| "You are SmolGPT-Fables. Write a vivid, complete fable from the user's canvas. " | |
| "Output only the finished story continuation. Begin with `### Scene 01:` and " | |
| "emit exactly the requested consecutive, zero-padded scene sections. A scene " | |
| "heading may include a short title after the colon. Use no other Markdown " | |
| "heading. Copy every required name, setting, and unusual detail verbatim. Make " | |
| "each character's described role, personality, and desire affect what they do. " | |
| "Write concrete action and dialogue instead of summarizing instructions. Keep " | |
| "each scene concise, make every scene change the situation, and resolve the " | |
| "ending target inside the final scene. Stop immediately after the final sentence; " | |
| "never add notes, analysis, a moral label, an ending section, or quoted canvas " | |
| "text. /no_think" | |
| ) | |
| _SCENE_COUNT_PATTERN = re.compile(r"(?m)^- Scene Count: ([0-9]+)$") | |
| _TARGET_SCENES_PATTERN = re.compile(r"(?m)^- Target scenes: ([0-9]+)$") | |
| _MIN_SCENE_COUNT = 1 | |
| _MAX_SCENE_COUNT = 6 | |
| SUPPORTED_PROMPT_FORMATS = frozenset( | |
| {CHAT_PROMPT_FORMAT, SMOLLM3_CHAT_PROMPT_FORMAT, RAW_PROMPT_FORMAT} | |
| ) | |
| CHAT_PROMPT_FORMATS = frozenset( | |
| {CHAT_PROMPT_FORMAT, SMOLLM3_CHAT_PROMPT_FORMAT} | |
| ) | |
| COMMON_ARTIFACT_FILES = ("config.json", "tokenizer.json") | |
| CUSTOM_CODE_FILES = ("configuration_smolgpt.py", "modeling_smolgpt.py") | |
| def _render_output_contract(scene_count: int) -> str: | |
| headings = ", ".join( | |
| f"`### Scene {index:02d}:`" for index in range(1, scene_count + 1) | |
| ) | |
| return CHAT_OUTPUT_CONTRACT_TEMPLATE.format( | |
| headings=headings, | |
| scene_count=scene_count, | |
| final_scene=scene_count, | |
| ) | |
| def prompt_contract_sha256() -> str: | |
| payload = { | |
| "format": CHAT_PROMPT_FORMAT, | |
| "system": CHAT_SYSTEM_PROMPT, | |
| "messages": ["system", "user", "assistant"], | |
| "assistant_only_loss": True, | |
| "prompt_transform": { | |
| "version": CHAT_PROMPT_TRANSFORM_VERSION, | |
| "scene_count_pattern": _SCENE_COUNT_PATTERN.pattern, | |
| "target_scenes_pattern": _TARGET_SCENES_PATTERN.pattern, | |
| "scene_count_range": [_MIN_SCENE_COUNT, _MAX_SCENE_COUNT], | |
| "story_boundary": CHAT_STORY_BOUNDARY, | |
| "rendered_output_contracts": { | |
| str(scene_count): _render_output_contract(scene_count) | |
| for scene_count in range(_MIN_SCENE_COUNT, _MAX_SCENE_COUNT + 1) | |
| }, | |
| }, | |
| } | |
| encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode() | |
| return hashlib.sha256(encoded).hexdigest() | |
| def smollm3_render_output_contract(scene_count: int) -> str: | |
| """Render the exact SmolLM3 v5 output contract used for SFT.""" | |
| if not _MIN_SCENE_COUNT <= scene_count <= _MAX_SCENE_COUNT: | |
| raise ValueError("v4 scene count must be between 1 and 6") | |
| headings = ", ".join( | |
| f"`### Scene {index:02d}:`" | |
| for index in range(1, scene_count + 1) | |
| ) | |
| minimum, maximum = SMOLLM3_SCENE_WORD_RANGE | |
| return ( | |
| "Output contract (follow exactly):\n" | |
| f"- Use these scene prefixes in order: {headings}.\n" | |
| f"- Write exactly {scene_count} scenes and {minimum}-{maximum} words per scene.\n" | |
| "- Use no heading except those scene headings.\n" | |
| "- Copy every item on `Must include` verbatim into the story.\n" | |
| "- Show the character-role details through decisions, action, or dialogue.\n" | |
| f"- Resolve the ending target in Scene {scene_count:02d} and stop.\n\n" | |
| ) | |
| def _canonical_smollm3_scene_count(prompt: str) -> int: | |
| scene_counts = _SCENE_COUNT_PATTERN.findall(prompt) | |
| target_counts = _TARGET_SCENES_PATTERN.findall(prompt) | |
| if len(scene_counts) != 1 or len(target_counts) != 1: | |
| raise ValueError("v4 prompt needs one Scene Count and one Target scenes line") | |
| scene_count = int(scene_counts[0]) | |
| target_count = int(target_counts[0]) | |
| if not _MIN_SCENE_COUNT <= scene_count <= _MAX_SCENE_COUNT: | |
| raise ValueError("v4 scene count must be between 1 and 6") | |
| if scene_count != target_count: | |
| raise ValueError("v4 Scene Count and Target scenes must match") | |
| return scene_count | |
| def smollm3_transform_prompt(prompt: str) -> str: | |
| """Apply the exact SmolLM3 v5 prompt transform used for SFT.""" | |
| scene_count = _canonical_smollm3_scene_count(prompt) | |
| if not prompt.endswith(SMOLLM3_STORY_BOUNDARY): | |
| raise ValueError("v4 prompt must end at the canonical Story boundary") | |
| return ( | |
| prompt[: -len(SMOLLM3_STORY_BOUNDARY)] | |
| + smollm3_render_output_contract(scene_count) | |
| + SMOLLM3_STORY_BOUNDARY | |
| ) | |
| def smollm3_chat_messages( | |
| prompt: str, | |
| completion: str | None = None, | |
| ) -> list[dict[str, str]]: | |
| messages = [ | |
| {"role": "system", "content": SMOLLM3_SYSTEM_PROMPT}, | |
| {"role": "user", "content": smollm3_transform_prompt(prompt)}, | |
| ] | |
| if completion is not None: | |
| messages.append({"role": "assistant", "content": completion}) | |
| return messages | |
| def smollm3_prompt_contract_sha256() -> str: | |
| """Hash the exact SmolLM3 v5 prompt contract used for SFT.""" | |
| payload = { | |
| "format": SMOLLM3_CHAT_PROMPT_FORMAT, | |
| "system": SMOLLM3_SYSTEM_PROMPT, | |
| "messages": ["system", "user", "assistant"], | |
| "assistant_only_loss": True, | |
| "thinking": False, | |
| "context_length": SMOLLM3_PRODUCT_CONTEXT_LENGTH, | |
| "transform_version": SMOLLM3_PROMPT_TRANSFORM_VERSION, | |
| "story_boundary": SMOLLM3_STORY_BOUNDARY, | |
| "scene_count_pattern": _SCENE_COUNT_PATTERN.pattern, | |
| "target_scenes_pattern": _TARGET_SCENES_PATTERN.pattern, | |
| "scene_word_range": list(SMOLLM3_SCENE_WORD_RANGE), | |
| "rendered_contracts": { | |
| str(count): smollm3_render_output_contract(count) | |
| for count in range(_MIN_SCENE_COUNT, _MAX_SCENE_COUNT + 1) | |
| }, | |
| } | |
| encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode() | |
| return hashlib.sha256(encoded).hexdigest() | |
| def raw_prompt_contract_sha256() -> str: | |
| payload = { | |
| "format": RAW_PROMPT_FORMAT, | |
| "messages": ["raw-markdown-prompt"], | |
| "bos_prefix": True, | |
| } | |
| encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode() | |
| return hashlib.sha256(encoded).hexdigest() | |
| def product_context_length_for_prompt_format(prompt_format: str) -> int | None: | |
| if prompt_format == CHAT_PROMPT_FORMAT: | |
| return PRODUCT_CONTEXT_LENGTH | |
| if prompt_format == SMOLLM3_CHAT_PROMPT_FORMAT: | |
| return SMOLLM3_PRODUCT_CONTEXT_LENGTH | |
| if prompt_format == RAW_PROMPT_FORMAT: | |
| return None | |
| raise ValueError(f"unsupported prompt format: {prompt_format}") | |
| def _has_model_weights(root: Path) -> bool: | |
| return (root / "model.safetensors").is_file() or ( | |
| root / "model.safetensors.index.json" | |
| ).is_file() | |
| def validate_transformers_artifact(root: Path) -> tuple[Mapping[str, Any], str]: | |
| """Validate either the legacy custom export or a standard Transformers LM.""" | |
| missing = [name for name in COMMON_ARTIFACT_FILES if not (root / name).is_file()] | |
| if missing: | |
| raise ValueError("model repository is missing: " + ", ".join(missing)) | |
| if not _has_model_weights(root): | |
| raise ValueError("model repository is missing Safetensors weights") | |
| try: | |
| config = json.loads((root / "config.json").read_text(encoding="utf-8")) | |
| except (OSError, json.JSONDecodeError) as exc: | |
| raise ValueError(f"could not read config.json: {exc}") from exc | |
| if not isinstance(config, Mapping): | |
| raise ValueError("config.json must contain a JSON object") | |
| auto_map = config.get("auto_map") | |
| if isinstance(auto_map, Mapping) and auto_map.get("AutoModelForCausalLM"): | |
| missing_code = [name for name in CUSTOM_CODE_FILES if not (root / name).is_file()] | |
| if missing_code: | |
| raise ValueError("custom model repository is missing: " + ", ".join(missing_code)) | |
| return config, "custom" | |
| architectures = config.get("architectures") | |
| if not isinstance(architectures, list) or not all( | |
| isinstance(value, str) and value for value in architectures | |
| ): | |
| raise ValueError("standard model config needs a non-empty architectures list") | |
| if not isinstance(config.get("model_type"), str) or not config["model_type"]: | |
| raise ValueError("standard model config needs model_type") | |
| return config, "standard" | |
| def _expected_prompt_contract_sha256(prompt_format: str) -> str: | |
| if prompt_format == CHAT_PROMPT_FORMAT: | |
| return prompt_contract_sha256() | |
| if prompt_format == SMOLLM3_CHAT_PROMPT_FORMAT: | |
| return smollm3_prompt_contract_sha256() | |
| if prompt_format == RAW_PROMPT_FORMAT: | |
| return raw_prompt_contract_sha256() | |
| raise ValueError(f"unsupported prompt format: {prompt_format}") | |
| def _manifest_prompt_binding( | |
| manifest: Mapping[str, Any], | |
| *, | |
| artifact_kind: str, | |
| ) -> tuple[str, str] | None: | |
| """Read legacy top-level or SmolLM3 nested bindings without ambiguity.""" | |
| top_format_present = "prompt_format" in manifest | |
| top_hash_present = "prompt_contract_sha256" in manifest | |
| nested_present = "prompt_contract" in manifest | |
| if not top_format_present and not top_hash_present and not nested_present: | |
| if artifact_kind == "custom": | |
| return None | |
| raise ValueError("training manifest has no supported prompt_format") | |
| if top_format_present != top_hash_present: | |
| raise ValueError("training manifest prompt binding is incomplete") | |
| bindings: list[tuple[str, str]] = [] | |
| if top_format_present: | |
| bindings.append( | |
| (manifest.get("prompt_format"), manifest.get("prompt_contract_sha256")) | |
| ) | |
| if nested_present: | |
| nested = manifest.get("prompt_contract") | |
| if ( | |
| not isinstance(nested, Mapping) | |
| or set(nested) != {"format", "sha256", "thinking"} | |
| or nested.get("thinking") is not False | |
| ): | |
| raise ValueError("training manifest nested prompt contract is invalid") | |
| bindings.append((nested.get("format"), nested.get("sha256"))) | |
| if any( | |
| not isinstance(value, str) | |
| or not value | |
| or not isinstance(digest, str) | |
| or re.fullmatch(r"[0-9a-f]{64}", digest) is None | |
| for value, digest in bindings | |
| ): | |
| raise ValueError("training manifest prompt binding is invalid") | |
| if any(binding != bindings[0] for binding in bindings[1:]): | |
| raise ValueError("training manifest prompt bindings conflict") | |
| return bindings[0] | |
| def prompt_format_for_artifact(root: Path, artifact_kind: str, tokenizer: Any) -> str: | |
| manifest_path = root / "training_manifest.json" | |
| if manifest_path.is_file(): | |
| manifest = json.loads(manifest_path.read_text(encoding="utf-8")) | |
| if not isinstance(manifest, Mapping): | |
| raise ValueError("training manifest must contain a JSON object") | |
| binding = _manifest_prompt_binding(manifest, artifact_kind=artifact_kind) | |
| if binding is None: | |
| # Checked-in pre-contract SmolGPT exports are unambiguously the raw | |
| # Markdown/BOS path. Any partially declared contract still fails. | |
| return RAW_PROMPT_FORMAT | |
| value, contract_hash = binding | |
| if value not in SUPPORTED_PROMPT_FORMATS: | |
| raise ValueError("training manifest has no supported prompt_format") | |
| if artifact_kind == "standard" and value not in CHAT_PROMPT_FORMATS: | |
| raise ValueError("standard product manifest must use the chat prompt format") | |
| if artifact_kind == "custom" and value == SMOLLM3_CHAT_PROMPT_FORMAT: | |
| raise ValueError("SmolLM3 chat-v5 requires a standard model artifact") | |
| expected_sha = _expected_prompt_contract_sha256(value) | |
| if contract_hash != expected_sha: | |
| raise ValueError("training manifest prompt contract hash does not match runtime") | |
| if value in CHAT_PROMPT_FORMATS and not getattr( | |
| tokenizer, "chat_template", None | |
| ): | |
| raise ValueError("chat prompt format requires a tokenizer chat_template") | |
| return str(value) | |
| if artifact_kind == "standard" and getattr(tokenizer, "chat_template", None): | |
| return CHAT_PROMPT_FORMAT | |
| return RAW_PROMPT_FORMAT | |
| def _canonical_scene_count(prompt: str) -> int: | |
| scene_counts = _SCENE_COUNT_PATTERN.findall(prompt) | |
| target_counts = _TARGET_SCENES_PATTERN.findall(prompt) | |
| if len(scene_counts) != 1: | |
| raise ValueError( | |
| "chat-v3 prompt must contain exactly one canonical '- Scene Count: N' line" | |
| ) | |
| if len(target_counts) != 1: | |
| raise ValueError( | |
| "chat-v3 prompt must contain exactly one canonical '- Target scenes: N' line" | |
| ) | |
| scene_count = int(scene_counts[0]) | |
| target_count = int(target_counts[0]) | |
| if not _MIN_SCENE_COUNT <= scene_count <= _MAX_SCENE_COUNT: | |
| raise ValueError("chat-v3 Scene Count must be between 1 and 6") | |
| if not _MIN_SCENE_COUNT <= target_count <= _MAX_SCENE_COUNT: | |
| raise ValueError("chat-v3 Target scenes must be between 1 and 6") | |
| if scene_count != target_count: | |
| raise ValueError("chat-v3 Scene Count and Target scenes must match") | |
| return scene_count | |
| def _transform_chat_prompt(prompt: str) -> str: | |
| scene_count = _canonical_scene_count(prompt) | |
| if not prompt.endswith(CHAT_STORY_BOUNDARY): | |
| raise ValueError("chat-v3 prompt must end at the canonical '## Story' boundary") | |
| return ( | |
| prompt[: -len(CHAT_STORY_BOUNDARY)] | |
| + _render_output_contract(scene_count) | |
| + CHAT_STORY_BOUNDARY | |
| ) | |
| def chat_messages(prompt: str, completion: str | None = None) -> list[dict[str, str]]: | |
| messages = [ | |
| {"role": "system", "content": CHAT_SYSTEM_PROMPT}, | |
| {"role": "user", "content": _transform_chat_prompt(prompt)}, | |
| ] | |
| if completion is not None: | |
| messages.append({"role": "assistant", "content": completion}) | |
| return messages | |
| def _chat_messages_for_format( | |
| prompt: str, | |
| completion: str | None, | |
| prompt_format: str, | |
| ) -> list[dict[str, str]]: | |
| if prompt_format == CHAT_PROMPT_FORMAT: | |
| return chat_messages(prompt, completion) | |
| if prompt_format == SMOLLM3_CHAT_PROMPT_FORMAT: | |
| return smollm3_chat_messages(prompt, completion) | |
| raise ValueError(f"unsupported chat prompt format: {prompt_format}") | |
| def _apply_runtime_chat_template( | |
| tokenizer: Any, | |
| messages: list[dict[str, str]], | |
| *, | |
| prompt_format: str, | |
| add_generation_prompt: bool, | |
| ) -> Any: | |
| kwargs: dict[str, Any] = { | |
| "add_generation_prompt": add_generation_prompt, | |
| "tokenize": True, | |
| } | |
| if prompt_format == SMOLLM3_CHAT_PROMPT_FORMAT: | |
| kwargs["enable_thinking"] = False | |
| return tokenizer.apply_chat_template(messages, **kwargs) | |
| def _flatten_token_ids(values: Any, *, context: str) -> list[int]: | |
| """Normalize chat-template outputs across supported Transformers versions.""" | |
| if isinstance(values, Mapping): | |
| if "input_ids" not in values: | |
| raise ValueError(f"{context} returned no input_ids") | |
| values = values["input_ids"] | |
| if hasattr(values, "tolist"): | |
| values = values.tolist() | |
| if not isinstance(values, Sequence) or isinstance(values, (str, bytes, bytearray)): | |
| raise ValueError(f"{context} returned unsupported token IDs") | |
| normalized = list(values) | |
| if normalized and isinstance(normalized[0], Sequence) and not isinstance( | |
| normalized[0], (str, bytes, bytearray) | |
| ): | |
| if len(normalized) != 1: | |
| raise ValueError(f"{context} returned more than one token sequence") | |
| normalized = list(normalized[0]) | |
| try: | |
| return [int(value) for value in normalized] | |
| except (TypeError, ValueError) as exc: | |
| raise ValueError(f"{context} returned non-integer token IDs") from exc | |
| def generation_prompt_ids(tokenizer: Any, prompt: str, prompt_format: str) -> list[int]: | |
| if prompt_format in CHAT_PROMPT_FORMATS: | |
| values = _apply_runtime_chat_template( | |
| tokenizer, | |
| _chat_messages_for_format(prompt, None, prompt_format), | |
| prompt_format=prompt_format, | |
| add_generation_prompt=True, | |
| ) | |
| context = ( | |
| "SmolLM3 v5 generation template" | |
| if prompt_format == SMOLLM3_CHAT_PROMPT_FORMAT | |
| else "chat generation template" | |
| ) | |
| return _flatten_token_ids(values, context=context) | |
| if prompt_format != RAW_PROMPT_FORMAT: | |
| raise ValueError(f"unsupported prompt format: {prompt_format}") | |
| return [ | |
| int(tokenizer.bos_token_id), | |
| *( | |
| int(value) | |
| for value in tokenizer.encode(prompt, add_special_tokens=False) | |
| ), | |
| ] | |
| def assistant_training_ids( | |
| tokenizer: Any, | |
| prompt: str, | |
| completion: str, | |
| *, | |
| max_length: int, | |
| prompt_format: str = CHAT_PROMPT_FORMAT, | |
| ) -> tuple[list[int], list[int]]: | |
| """Create one chat sequence with loss masked through the assistant header.""" | |
| if prompt_format not in CHAT_PROMPT_FORMATS: | |
| raise ValueError("assistant training requires a supported chat prompt format") | |
| prefix = generation_prompt_ids(tokenizer, prompt, prompt_format) | |
| full = _flatten_token_ids( | |
| _apply_runtime_chat_template( | |
| tokenizer, | |
| _chat_messages_for_format(prompt, completion, prompt_format), | |
| prompt_format=prompt_format, | |
| add_generation_prompt=False, | |
| ), | |
| context=( | |
| "SmolLM3 v5 training template" | |
| if prompt_format == SMOLLM3_CHAT_PROMPT_FORMAT | |
| else "chat training template" | |
| ), | |
| ) | |
| if full[: len(prefix)] != prefix: | |
| raise ValueError("chat template assistant prefix is not stable") | |
| if len(full) > max_length: | |
| raise ValueError( | |
| f"chat-formatted example has {len(full)} tokens; maximum is {max_length}" | |
| ) | |
| if len(full) <= len(prefix): | |
| raise ValueError("chat-formatted example has no assistant completion tokens") | |
| labels = [-100] * len(prefix) + full[len(prefix) :] | |
| return full, labels | |
| def aggregate_sha256(paths: Sequence[Path], root: Path) -> str: | |
| records = [] | |
| for path in sorted(paths): | |
| records.append( | |
| { | |
| "path": path.relative_to(root).as_posix(), | |
| "sha256": hashlib.sha256(path.read_bytes()).hexdigest(), | |
| } | |
| ) | |
| payload = json.dumps(records, sort_keys=True, separators=(",", ":")).encode() | |
| return hashlib.sha256(payload).hexdigest() | |