Dolphin3.0-CoreML / scripts /generate.py
ales27pm's picture
Make the Core ML release self-contained and runnable
95671cf verified
Raw
History Blame Contribute Delete
5.85 kB
#!/usr/bin/env python3
# /// script
# requires-python = ">=3.11,<3.12"
# dependencies = [
# "coremltools==8.0",
# "jinja2==3.1.5",
# "numpy==1.26.4",
# "transformers==4.47.1",
# ]
# ///
"""Run bounded text generation with the stateful Dolphin Core ML package."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Iterable
import coremltools as ct
import numpy as np
from transformers import AutoTokenizer
DEFAULT_MODEL = "Dolphin3.0-Llama3.2-3B-stateful-int4.mlpackage"
DEFAULT_TOKENIZER = "ales27pm/Dolphin3.0-CoreML"
STOP_TOKEN_IDS = frozenset((128256, 128001, 128008, 128009))
def causal_mask(query_length: int, end_step: int) -> np.ndarray:
if query_length < 1 or end_step < query_length:
raise ValueError("Expected 1 <= query_length <= end_step")
past_length = end_step - query_length
columns = np.arange(end_step)[None, :]
rows = past_length + np.arange(query_length)[:, None]
return np.where(columns <= rows, 0.0, -65504.0).astype(np.float16)[
None, None, :, :
]
def sample_token(
logits: np.ndarray,
*,
temperature: float,
top_p: float,
rng: np.random.Generator,
) -> int:
scores = logits[0, -1].astype(np.float32)
if not np.isfinite(scores).all():
raise RuntimeError("Core ML returned non-finite logits")
if temperature <= 0:
return int(np.argmax(scores))
scores /= temperature
scores -= np.max(scores)
probabilities = np.exp(scores)
probabilities /= probabilities.sum()
order = np.argsort(probabilities)[::-1]
ordered = probabilities[order]
# Keep the first token whose inclusion reaches or crosses the requested
# probability mass. Subtracting the current probability makes the test
# equivalent to shifting the cumulative mask one position to the right.
keep = np.cumsum(ordered) - ordered < top_p
selected = order[keep]
selected_probabilities = probabilities[selected]
selected_probabilities /= selected_probabilities.sum()
return int(rng.choice(selected, p=selected_probabilities))
def stop_ids(tokenizer_eos: int | Iterable[int] | None) -> frozenset[int]:
values = set(STOP_TOKEN_IDS)
if isinstance(tokenizer_eos, int):
values.add(tokenizer_eos)
elif tokenizer_eos is not None:
values.update(int(item) for item in tokenizer_eos)
return frozenset(values)
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("prompt")
parser.add_argument("--model", default=DEFAULT_MODEL)
parser.add_argument("--tokenizer", default=DEFAULT_TOKENIZER)
parser.add_argument(
"--system",
default="You are Dolphin, created by Eric Hartford. You are a helpful assistant.",
)
parser.add_argument("--max-new-tokens", type=int, default=64)
parser.add_argument("--temperature", type=float, default=0.0)
parser.add_argument("--top-p", type=float, default=0.9)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument(
"--compute-units",
choices=("all", "cpu_and_gpu", "cpu_only", "cpu_and_ne"),
default="cpu_and_gpu",
)
args = parser.parse_args()
if not 0 < args.top_p <= 1:
parser.error("--top-p must be in (0, 1]")
if args.max_new_tokens < 1:
parser.error("--max-new-tokens must be positive")
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer, revision="main")
messages = [
{"role": "system", "content": args.system},
{"role": "user", "content": args.prompt},
]
prompt_ids = tokenizer.apply_chat_template(
messages, add_generation_prompt=True, return_tensors="np"
).astype(np.int32)
compute_units = {
"all": ct.ComputeUnit.ALL,
"cpu_and_gpu": ct.ComputeUnit.CPU_AND_GPU,
"cpu_only": ct.ComputeUnit.CPU_ONLY,
"cpu_and_ne": ct.ComputeUnit.CPU_AND_NE,
}[args.compute_units]
model = ct.models.MLModel(args.model, compute_units=compute_units)
metadata = model.user_defined_metadata
max_context = int(
metadata.get("com.ales27pm.dolphin.max_context_length", "2048")
)
max_query = int(metadata.get("com.ales27pm.dolphin.max_query_length", "512"))
if prompt_ids.shape[-1] > max_query:
raise ValueError(
f"Prompt has {prompt_ids.shape[-1]} tokens; model prefill limit is {max_query}"
)
if prompt_ids.shape[-1] + args.max_new_tokens > max_context:
raise ValueError(
"Prompt plus requested output exceeds the model's "
f"{max_context}-token state capacity"
)
state = model.make_state()
rng = np.random.default_rng(args.seed)
generated: list[int] = []
query = prompt_ids
end_step = prompt_ids.shape[-1]
eos_ids = stop_ids(tokenizer.eos_token_id)
for _ in range(args.max_new_tokens):
result = model.predict(
{
"inputIds": query,
"causalMask": causal_mask(query.shape[-1], end_step),
},
state=state,
)
token = sample_token(
result["logits"],
temperature=args.temperature,
top_p=args.top_p,
rng=rng,
)
if token in eos_ids:
break
generated.append(token)
query = np.array([[token]], dtype=np.int32)
end_step += 1
text = tokenizer.decode(generated, skip_special_tokens=True)
print(text)
print(
json.dumps(
{
"prompt_tokens": int(prompt_ids.shape[-1]),
"generated_tokens": len(generated),
"stop_token_ids": sorted(eos_ids),
},
sort_keys=True,
)
)
return 0
if __name__ == "__main__":
raise SystemExit(main())