Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 57 additions & 17 deletions ollama/_types.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,19 @@
import contextlib
import json
import warnings
from base64 import b64decode, b64encode
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, List, Mapping, Optional, Sequence, Union

from pydantic import (
BaseModel,
BeforeValidator,
ByteSize,
ConfigDict,
Field,
model_serializer,
model_validator,
)
from pydantic.json_schema import JsonSchemaValue
from typing_extensions import Annotated, Literal
Expand Down Expand Up @@ -101,41 +104,78 @@ def get(self, key: str, default: Any = None) -> Any:
return getattr(self, key) if hasattr(self, key) else default


# Warn on no longer supported options
_UNSUPPORTED_OPTIONS = frozenset(
{
'embedding_only',
'f16_kv',
'logits_all',
'low_vram',
'mirostat',
'mirostat_eta',
'mirostat_tau',
'numa',
'penalize_newline',
'tfs_z',
'typical_p',
'use_mlock',
'vocab_only',
}
)


def _drop_unsupported_options(options: Any) -> Any:
if not isinstance(options, Mapping):
return options
Comment on lines +128 to +129

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This still lets typical_p through when someone uses Options().model_copy(update={'typical_p': 0.5}). I checked the request body and it’s sent without a warning, so the server would still reject it. Could we strip these options before sending the request and add a test for this case?

unsupported = [key for key in options if key in _UNSUPPORTED_OPTIONS]
if not unsupported:
return options
for key in unsupported:
warnings.warn(f'option {key!r} is no longer supported and was ignored', FutureWarning, stacklevel=2)
return {key: value for key, value in options.items() if key not in _UNSUPPORTED_OPTIONS}


class Options(SubscriptableBaseModel):
# Unknown options pass through so new server options work before the client is updated.
model_config = ConfigDict(extra='allow')

# load time options
numa: Optional[bool] = None
num_ctx: Optional[int] = None
num_batch: Optional[int] = None
num_gpu: Optional[int] = None
main_gpu: Optional[int] = None
low_vram: Optional[bool] = None
f16_kv: Optional[bool] = None
logits_all: Optional[bool] = None
vocab_only: Optional[bool] = None
use_mmap: Optional[bool] = None
use_mlock: Optional[bool] = None
embedding_only: Optional[bool] = None
num_thread: Optional[int] = None
draft_num_predict: Optional[int] = None

# runtime options
num_keep: Optional[int] = None
seed: Optional[int] = None
num_predict: Optional[int] = None
top_k: Optional[int] = None
top_p: Optional[float] = None
tfs_z: Optional[float] = None
typical_p: Optional[float] = None

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could we keep these fields around and just leave them out of requests? Existing code that reads options.typical_p now crashes, even though passing it into Options still works with a warning. A test that checks reading the old fields would help here.

min_p: Optional[float] = None
repeat_last_n: Optional[int] = None
temperature: Optional[float] = None
repeat_penalty: Optional[float] = None
presence_penalty: Optional[float] = None
frequency_penalty: Optional[float] = None
mirostat: Optional[int] = None
mirostat_tau: Optional[float] = None
mirostat_eta: Optional[float] = None
penalize_newline: Optional[bool] = None
stop: Optional[Sequence[str]] = None

@model_validator(mode='before')
@classmethod
def drop_unsupported(cls, data: Any) -> Any:
return _drop_unsupported_options(data)

def __setattr__(self, name: str, value: Any) -> None:
if name in _UNSUPPORTED_OPTIONS:
warnings.warn(f'option {name!r} is no longer supported and was ignored', FutureWarning, stacklevel=2)
return
super().__setattr__(name, value)


_RequestOptions = Annotated[Optional[Union[Mapping[str, Any], Options]], BeforeValidator(_drop_unsupported_options)]


class BaseRequest(SubscriptableBaseModel):
model: Annotated[str, Field(min_length=1)]
Expand All @@ -148,7 +188,7 @@ class BaseStreamableRequest(BaseRequest):


class BaseGenerateRequest(BaseStreamableRequest):
options: Optional[Union[Mapping[str, Any], Options]] = None
options: _RequestOptions = None
'Options to use for the request.'

format: Optional[Union[Literal['', 'json'], JsonSchemaValue]] = None
Expand Down Expand Up @@ -429,7 +469,7 @@ class EmbedRequest(BaseRequest):
truncate: Optional[bool] = None
'Truncate the input to the maximum token length.'

options: Optional[Union[Mapping[str, Any], Options]] = None
options: _RequestOptions = None
'Options to use for the request.'

keep_alive: Optional[Union[float, str]] = None
Expand All @@ -451,7 +491,7 @@ class EmbeddingsRequest(BaseRequest):
prompt: Optional[str] = None
'Prompt to generate embeddings from.'

options: Optional[Union[Mapping[str, Any], Options]] = None
options: _RequestOptions = None
'Options to use for the request.'

keep_alive: Optional[Union[float, str]] = None
Expand Down Expand Up @@ -502,7 +542,7 @@ def serialize_model(self, nxt):
template: Optional[str] = None
license: Optional[Union[str, List[str]]] = None
system: Optional[str] = None
parameters: Optional[Union[Mapping[str, Any], Options]] = None
parameters: _RequestOptions = None
messages: Optional[Sequence[Union[Mapping[str, Any], Message]]] = None


Expand Down
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -37,8 +37,9 @@ dependencies = [ 'ruff>=0.9.1' ]
config-path = 'none'

[tool.ruff]
line-length = 320
line-length = 320
indent-width = 2
extend-exclude = ['*.md']

[tool.ruff.format]
quote-style = 'single'
Expand Down
58 changes: 58 additions & 0 deletions tests/test_options.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
import warnings

import pytest
from pytest_httpserver import HTTPServer

from ollama._client import Client
from ollama._types import ChatRequest, CreateRequest, Options


def test_options_drop_unsupported_keywords():
with pytest.warns(FutureWarning, match='typical_p'):
options = Options(typical_p=0.5, temperature=0.1)
assert options.model_dump(exclude_none=True) == {'temperature': 0.1}


def test_options_drop_unsupported_item_assignment():
options = Options(temperature=0.1)
with pytest.warns(FutureWarning, match='mirostat'):
options['mirostat'] = 2
assert options.model_dump(exclude_none=True) == {'temperature': 0.1}


def test_options_keep_unknown_keys():
with warnings.catch_warnings():
warnings.simplefilter('error')
options = Options(future_option=1, min_p=0.05)
assert options.model_dump(exclude_none=True) == {'future_option': 1, 'min_p': 0.05}


def test_request_drops_unsupported_options_mapping():
with pytest.warns(FutureWarning, match='typical_p'):
request = ChatRequest(model='dummy', messages=[], options={'typical_p': 0.5, 'num_ctx': 8})
assert request.model_dump(exclude_none=True)['options'] == {'num_ctx': 8}


def test_create_request_drops_unsupported_parameters():
with pytest.warns(FutureWarning, match='penalize_newline'):
request = CreateRequest(model='dummy', parameters={'penalize_newline': True, 'pi': 3.14})
assert request.model_dump(exclude_none=True)['parameters'] == {'pi': 3.14}


def test_client_chat_drops_unsupported_options(httpserver: HTTPServer):
httpserver.expect_ordered_request(
'/api/chat',
method='POST',
json={
'model': 'dummy',
'messages': [{'role': 'user', 'content': 'Hi'}],
'tools': [],
'stream': False,
'options': {'temperature': 0.0},
},
).respond_with_json({'model': 'dummy', 'message': {'role': 'assistant', 'content': 'Hello'}})

client = Client(httpserver.url_for('/'))
with pytest.warns(FutureWarning, match='typical_p'):
response = client.chat('dummy', messages=[{'role': 'user', 'content': 'Hi'}], options={'typical_p': 0.5, 'temperature': 0.0})
assert response['message']['content'] == 'Hello'
Loading