Spaces:
Sleeping
Sleeping
File size: 20,397 Bytes
095660c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 | """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()
|