From 71bbe5f25b463eecb953f292fc2aa7fad1f9f2e7 Mon Sep 17 00:00:00 2001 From: Ryuichi Leo Takashige Date: Thu, 22 Jan 2026 14:51:12 +0000 Subject: [PATCH] Review and extract logprob stuff from alexcheema/uncertainty-visualization --- src/exo/master/api.py | 35 ++++++++++ src/exo/shared/types/api.py | 2 + src/exo/shared/types/chunks.py | 4 +- .../shared/types/worker/runner_response.py | 8 +-- .../worker/engines/mlx/generator/generate.py | 65 ++++++++++++++++++- src/exo/worker/runner/runner.py | 2 + 6 files changed, 107 insertions(+), 9 deletions(-) diff --git a/src/exo/master/api.py b/src/exo/master/api.py index 5c0a46b1..c1c26314 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -51,6 +51,8 @@ from exo.shared.types.api import ( ImageGenerationTaskParams, ImageListItem, ImageListResponse, + Logprobs, + LogprobsContentItem, ModelList, ModelListModel, PlaceInstanceParams, @@ -100,9 +102,27 @@ def _format_to_content_type(image_format: Literal["png", "jpeg", "webp"] | None) return f"image/{image_format or 'png'}" +def _build_logprobs(chunk: TokenChunk) -> Logprobs: + """Convert flat logprob fields to OpenAI Logprobs format.""" + return Logprobs( + content=[ + LogprobsContentItem( + token=chunk.text, + logprob=chunk.logprob if chunk.logprob is not None else 0.0, + bytes=list(chunk.text.encode("utf-8")), + top_logprobs=chunk.top_logprobs or [], + ) + ] + ) + + def chunk_to_response( chunk: TokenChunk | ToolCallChunk, command_id: CommandId ) -> ChatCompletionResponse: + logprobs: Logprobs | None = None + if isinstance(chunk, TokenChunk) and chunk.logprob is not None: + logprobs = _build_logprobs(chunk) + return ChatCompletionResponse( id=command_id, created=int(time.time()), @@ -123,6 +143,7 @@ def chunk_to_response( for i, tool in enumerate(chunk.tool_calls) ], ), + logprobs=logprobs, finish_reason=chunk.finish_reason, ) ], @@ -527,6 +548,7 @@ class API: text_parts: list[str] = [] tool_calls: list[ToolCall] = [] + logprobs_items: list[LogprobsContentItem] = [] model: str | None = None finish_reason: FinishReason | None = None @@ -542,6 +564,14 @@ class API: if isinstance(chunk, TokenChunk): text_parts.append(chunk.text) + if chunk.logprob is not None: + lp = _build_logprobs(chunk) + if lp.content: + if len(lp.content) != 1: + logger.warning( + f"Expected 1 logprobs content item per chunk, got {len(lp.content)}" + ) + logprobs_items.append(lp.content[0]) if isinstance(chunk, ToolCallChunk): tool_calls.extend( @@ -559,6 +589,10 @@ class API: combined_text = "".join(text_parts) assert model is not None + logprobs: Logprobs | None = None + if logprobs_items: + logprobs = Logprobs(content=logprobs_items) + return ChatCompletionResponse( id=command_id, created=int(time.time()), @@ -571,6 +605,7 @@ class API: content=combined_text, tool_calls=tool_calls, ), + logprobs=logprobs, finish_reason=finish_reason, ) ], diff --git a/src/exo/shared/types/api.py b/src/exo/shared/types/api.py index 3a977817..3b70395f 100644 --- a/src/exo/shared/types/api.py +++ b/src/exo/shared/types/api.py @@ -97,6 +97,8 @@ class LogprobsContentItem(BaseModel): class Logprobs(BaseModel): content: list[LogprobsContentItem] | None = None + # This will always be null for open source models, but exists for OpenAI API + refusal: list[LogprobsContentItem] | None = None class PromptTokensDetails(BaseModel): diff --git a/src/exo/shared/types/chunks.py b/src/exo/shared/types/chunks.py index 235ef70d..c776d7d7 100644 --- a/src/exo/shared/types/chunks.py +++ b/src/exo/shared/types/chunks.py @@ -2,7 +2,7 @@ from collections.abc import Generator from typing import Any, Literal from exo.shared.models.model_cards import ModelId -from exo.shared.types.api import GenerationStats, ImageGenerationStats +from exo.shared.types.api import GenerationStats, ImageGenerationStats, TopLogprobItem from exo.utils.pydantic_ext import TaggedModel from .api import FinishReason @@ -17,6 +17,8 @@ class BaseChunk(TaggedModel): class TokenChunk(BaseChunk): text: str token_id: int + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None finish_reason: Literal["stop", "length", "content_filter"] | None = None stats: GenerationStats | None = None diff --git a/src/exo/shared/types/worker/runner_response.py b/src/exo/shared/types/worker/runner_response.py index 8d695ab0..a1f29ebd 100644 --- a/src/exo/shared/types/worker/runner_response.py +++ b/src/exo/shared/types/worker/runner_response.py @@ -6,6 +6,7 @@ from exo.shared.types.api import ( GenerationStats, ImageGenerationStats, ToolCallItem, + TopLogprobItem, ) from exo.utils.pydantic_ext import TaggedModel @@ -14,14 +15,11 @@ class BaseRunnerResponse(TaggedModel): pass -class TokenizedResponse(BaseRunnerResponse): - prompt_tokens: int - - class GenerationResponse(BaseRunnerResponse): text: str token: int - # logprobs: list[float] | None = None # too big. we can change to be top-k + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None finish_reason: FinishReason | None = None stats: GenerationStats | None = None diff --git a/src/exo/worker/engines/mlx/generator/generate.py b/src/exo/worker/engines/mlx/generator/generate.py index b4ea2a6c..755c8cbf 100644 --- a/src/exo/worker/engines/mlx/generator/generate.py +++ b/src/exo/worker/engines/mlx/generator/generate.py @@ -12,12 +12,11 @@ from exo.shared.types.api import ( ChatCompletionMessage, FinishReason, GenerationStats, + TopLogprobItem, ) from exo.shared.types.memory import Memory from exo.shared.types.tasks import ChatCompletionTaskParams -from exo.shared.types.worker.runner_response import ( - GenerationResponse, -) +from exo.shared.types.worker.runner_response import GenerationResponse from exo.worker.engines.mlx import Model from exo.worker.engines.mlx.constants import KV_BITS, KV_GROUP_SIZE, MAX_TOKENS from exo.worker.engines.mlx.utils_mlx import ( @@ -115,6 +114,49 @@ def eos_ids_from_tokenizer(tokenizer: TokenizerWrapper) -> list[int]: return eos +def extract_top_logprobs( + logprobs_array: mx.array, + selected_token: int, + tokenizer: TokenizerWrapper, + top_k: int, +) -> tuple[float, list[TopLogprobItem]]: + """Extract the selected token's logprob and top-k alternatives. + + Returns: + Tuple of (selected_token_logprob, list of TopLogprobItem) + """ + selected_logprob = float(logprobs_array[selected_token].item()) + + if top_k == 0: + return selected_logprob, [] + + vocab_size = logprobs_array.shape[0] + k = min(top_k, vocab_size) + top_indices = mx.argpartition(-logprobs_array, kth=k - 1)[:k] + + top_logprobs_values = logprobs_array[top_indices] + sorted_order = mx.argsort(-top_logprobs_values) + top_indices = top_indices[sorted_order] + + mx.eval(top_indices) + + top_logprob_items: list[TopLogprobItem] = [] + indices_list: list[int] = cast(list[int], top_indices.tolist()) + for token_id in indices_list: + logprob_value = float(logprobs_array[token_id].item()) + token_str = tokenizer.decode([token_id]) + + top_logprob_items.append( + TopLogprobItem( + token=token_str, + logprob=logprob_value, + bytes=list(token_str.encode("utf-8")), + ) + ) + + return selected_logprob, top_logprob_items + + def mlx_generate( model: Model, tokenizer: TokenizerWrapper, @@ -144,6 +186,10 @@ def mlx_generate( top_p=task.top_p if task.top_p is not None else 1.0, ) + # Determine if we need logprobs + should_extract_logprobs = task.logprobs is True + top_k = task.top_logprobs if task.top_logprobs is not None else 0 + max_tokens = task.max_tokens or MAX_TOKENS for out in stream_generate( model=model, @@ -177,9 +223,22 @@ def mlx_generate( f"Model generated unexpected finish_reason: {out.finish_reason}" ) + # Extract logprobs if requested + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None + if should_extract_logprobs: + logprob, top_logprobs = extract_top_logprobs( + logprobs_array=out.logprobs, + selected_token=out.token, + tokenizer=tokenizer, + top_k=top_k, + ) + yield GenerationResponse( text=out.text, token=out.token, + logprob=logprob, + top_logprobs=top_logprobs, finish_reason=cast(FinishReason | None, out.finish_reason), stats=stats, ) diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index faad3ca9..b539d59e 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -297,6 +297,8 @@ def main( model=shard_metadata.model_card.model_id, text=response.text, token_id=response.token, + logprob=response.logprob, + top_logprobs=response.top_logprobs, finish_reason=response.finish_reason, stats=response.stats, ),