Skip to content
49 changes: 49 additions & 0 deletions examples/high_level_api/high_level_api_prefill.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,49 @@
"""Decision one token without sampling
"""
import argparse
import numpy as np
from llama_cpp import Llama

parser = argparse.ArgumentParser()
parser.add_argument("-m", "--model", type=str, default="../models/7B/ggml-model.bin")
args = parser.parse_args()

llm = Llama(
model_path=args.model,
n_ctx=2048,
)

prompt = """\
Question: Is the following statement true?

2 + 2 = 4

Answer with exactly Y or N.
Answer:
"""

result = llm.prefill(prompt)

y_tokens = llm.tokenize(b"Y", add_bos=False, special=False)
n_tokens = llm.tokenize(b"N", add_bos=False, special=False)

assert len(y_tokens) == 1
assert len(n_tokens) == 1

y_token = y_tokens[0]
n_token = n_tokens[0]

# pull Y/N logit from vocab
y_logit = float(result.logits[y_token])
n_logit = float(result.logits[n_token])

print("Y logit:", y_logit)
print("N logit:", n_logit)

# softmax from Y/N only
candidate_logits = np.array([y_logit, n_logit])
candidate_probs = np.exp(candidate_logits - candidate_logits.max())
candidate_probs /= candidate_probs.sum()

print("Y:", candidate_probs[0])
print("N:", candidate_probs[1])
96 changes: 96 additions & 0 deletions examples/high_level_api/high_level_api_prefill_vision.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
"""Decision one token without sampling (with vision).
Gemma4 model only.
"""
import numpy as np
import argparse
from pathlib import Path
from llama_cpp import Llama
from llama_cpp.llama_chat_format import Gemma4ChatHandler

BASE_DIR = Path(__file__).resolve().parent
IMAGE_PATH = BASE_DIR / "media" / "apple.jpg"

parser = argparse.ArgumentParser()
parser.add_argument("-m", "--model", type=str)
parser.add_argument("-mm", "--mmproj", type=str)
parser.add_argument("-i", "--image", type=str, default=IMAGE_PATH)
args = parser.parse_args()

# Use gemma4 handler
chat_handler = Gemma4ChatHandler(
clip_model_path=args.mmproj,
enable_thinking=False,
)

llm = Llama(
model_path=args.model,
chat_handler=chat_handler,
n_ctx=4096,
)

messages = [
{
"role": "system",
"content": "Answer with exactly Y or N."
},
{
"role": "user",
"content": [
{
"type": "image_url",
"image_url": {
"url": args.image.as_uri(),
},
},
{
"type": "text",
"text": "Is there any visible written word in this image? "
"Answer with exactly one character: Y or N."
},
],
},
]

result = llm.create_chat_prefill(
messages=messages,
)

result_chat = llm.create_chat_completion(
messages=messages,
)

print(result_chat)

print("Top 20 tokens: ")
top_k = 20
top_ids = np.argsort(result.logits)[-top_k:][::-1]

for token_id in top_ids:
token = llm.detokenize([int(token_id)])
print(
repr(token),
float(result.logits[token_id]),
)

y_tokens = llm.tokenize(b"Y", add_bos=False, special=False)
n_tokens = llm.tokenize(b"N", add_bos=False, special=False)

assert len(y_tokens) == 1
assert len(n_tokens) == 1

y_token = y_tokens[0]
n_token = n_tokens[0]

y_logit = float(result.logits[y_token])
n_logit = float(result.logits[n_token])

print("Y logit:", y_logit)
print("N logit:", n_logit)

# Softmax in Y/N
candidate_logits = np.array([y_logit, n_logit], dtype=np.float32)
candidate_probs = np.exp(candidate_logits - candidate_logits.max())
candidate_probs /= candidate_probs.sum()

print("P(Y | {Y,N}) =", float(candidate_probs[0]))
print("P(N | {Y,N}) =", float(candidate_probs[1]))
Binary file added examples/high_level_api/media/apple.jpg
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
94 changes: 88 additions & 6 deletions llama_cpp/llama.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@
import llama_cpp.llama_cpp as llama_cpp_lib
import llama_cpp.llama_chat_format as llama_chat_format
import llama_cpp.llama_multimodal as llama_multimodal
from .llama_chat_format import PrefillResult

from llama_cpp.llama_speculative import (
LlamaDraftModel,
Expand Down Expand Up @@ -1813,6 +1814,44 @@ def eval(
logits_view = np.ctypeslib.as_array(logits_ptr, shape=(self._n_vocab,))
self.scores[0, :] = logits_view

def prefill(
self,
prompt: Union[str, Sequence[int]],
*,
reset: bool = True,
add_bos: bool = True,
special: bool = True,
active_loras: Optional[List[Dict[str, Union[str, float]]]] = None,
control_vector: Optional[Dict[str, Any]] = None,
) -> PrefillResult:
"""Evaluate a text prompt and return its final next-token logits."""
tokens = (
self.tokenize(prompt.encode("utf-8"), add_bos=add_bos, special=special)
if isinstance(prompt, str)
else list(prompt)
)
if not tokens:
raise ValueError("Prefill requires at least one token")
if reset:
self.reset()

self.eval(
tokens,
active_loras=active_loras,
control_vector=control_vector,
copy_logits=True,
)

logits = (
self.scores[self.n_tokens - 1]
if self._logits_all
else self.scores[0]
)
return PrefillResult(
n_tokens=len(tokens),
logits=logits
)

# Helper method: Convert dict logit_bias to List[llama_logit_bias]
def _convert_logit_bias(self, logit_bias: Optional[Dict[int, float]]) -> List[llama_cpp_lib.llama_logit_bias]:
if not logit_bias:
Expand Down Expand Up @@ -4361,6 +4400,15 @@ def __call__(
presence_penalty=presence_penalty,
)

def _get_chat_completion_handler(
self,
) -> llama_chat_format.LlamaChatCompletionHandler:
return (
self.chat_handler
or self._chat_handlers.get(self.chat_format)
or llama_chat_format.get_chat_completion_handler(self.chat_format)
)

def create_chat_completion(
self,
messages: List[ChatCompletionRequestMessage],
Expand Down Expand Up @@ -4411,6 +4459,7 @@ def create_chat_completion(
top_logprobs: Optional[int] = None,
assistant_prefill: bool = False,
add_generation_prompt: bool = True,
chat_template_kwargs: Optional[Dict[str, Any]] = None,
# Reasoning Budget Params
reasoning_budget: int = -1,
reasoning_start: str = "<think>",
Expand Down Expand Up @@ -4457,7 +4506,7 @@ def create_chat_completion(
dry_base`: Set the DRY repetition penalty base value. Default: `1.75`
dry_allowed_length: Tokens that extend repetition beyond this receive exponentially increasing penalty: multiplier * base ^ (length of repeating sequence before token - allowed length). Default: `2`
dry_penalty_last_n: How many tokens to scan for repetitions. Default: `64`; `0` disables scanning and `-1` uses the context size.
dry_seq_breakers: Specify an array of sequence breakers for DRY sampling. Only a JSON array of strings is accepted. Default: `['\n', ':', '"', '*']`
dry_seq_breakers: Specify an array of sequence breakers for DRY sampling. Only a JSON array of strings is accepted. Default: `['\\n', ':', '"', '*']`
adaptive-target: Adaptive-p: select tokens near this probability (valid range 0.0 to 1.0; negative = disabled) (default: %.2f) [(more info)](https://github.com/ggml-org/llama.cpp/pull/17927)
adaptive-decay: Adaptive-p: decay rate for target adaptation over time. lower values are more reactive, higher values are more stable. (valid range 0.0 to 0.99) (default: %.2f)
use_infill: Determines whether to activate the specialized fill-in-the-middle sampler that consolidates probabilities of tokens sharing common prefixes to ensure the generated text coherently bridges the gap between the prefix and suffix.
Expand All @@ -4466,6 +4515,7 @@ def create_chat_completion(
logits_processor: A list of logits processors to use.
grammar: A grammar to use.
grammar_lazy: If True, enables lazy evaluation.
chat_template_kwargs: Optional keyword arguments passed to the Jinja chat template at render time. These values override matching handler-level template defaults for the current request only.
reasoning_budget: Token budget for the first visible reasoning block.
-1 disables the sampler, 0 forces an immediate end after reasoning starts,
and N > 0 allows at most N generated tokens inside the block.
Expand All @@ -4489,11 +4539,7 @@ def create_chat_completion(
if presence_penalty is not None and present_penalty == 0.0:
present_penalty = presence_penalty

handler = (
self.chat_handler
or self._chat_handlers.get(self.chat_format)
or llama_chat_format.get_chat_completion_handler(self.chat_format)
)
handler = self._get_chat_completion_handler()
return handler(
llama=self,
messages=messages,
Expand Down Expand Up @@ -4543,6 +4589,7 @@ def create_chat_completion(
control_vector=control_vector,
assistant_prefill=assistant_prefill,
add_generation_prompt=add_generation_prompt,
chat_template_kwargs=chat_template_kwargs,
reasoning_budget=reasoning_budget,
reasoning_start=reasoning_start,
reasoning_end=reasoning_end,
Expand All @@ -4551,6 +4598,41 @@ def create_chat_completion(
reasoning_start_max_tokens=reasoning_start_max_tokens,
)

def create_chat_prefill(
self,
messages: List[ChatCompletionRequestMessage],
functions: Optional[List[ChatCompletionFunction]] = None,
function_call: Optional[ChatCompletionRequestFunctionCall] = None,
tools: Optional[List[ChatCompletionTool]] = None,
tool_choice: Optional[ChatCompletionToolChoiceOption] = None,
add_generation_prompt: bool = True,
assistant_prefill: bool = False,
chat_template_kwargs: Optional[Dict[str, Any]] = None,
) -> PrefillResult:
"""Prefill a chat prompt through its handler without generating a token."""
handler = self._get_chat_completion_handler()
prefill = getattr(handler, "prefill", None)
if not callable(prefill):
raise NotImplementedError(
"The selected chat handler does not support prefill"
)

prefill_kwargs: Dict[str, Any] = {
"llama": self,
"messages": messages,
"functions": functions,
"function_call": function_call,
"tools": tools,
"tool_choice": tool_choice,
"add_generation_prompt": add_generation_prompt,
"assistant_prefill": assistant_prefill,
}
# For compatibility
if chat_template_kwargs is not None:
prefill_kwargs["chat_template_kwargs"] = chat_template_kwargs

return prefill(**prefill_kwargs)

def create_chat_completion_openai_v1(
self,
*args: Any,
Expand Down
Loading
Loading