Skip to content
Merged
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
4 changes: 2 additions & 2 deletions src/art/tau_bench/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,8 +18,8 @@

def _default_limits() -> httpx.Limits:
return httpx.Limits(
max_connections=512,
max_keepalive_connections=512,
max_connections=100_000,
max_keepalive_connections=100_000,
keepalive_expiry=60.0,
)

Expand Down
2 changes: 1 addition & 1 deletion src/art/tau_bench/rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
openai_clients: dict[tuple[str, str], AsyncOpenAI] = {}
CONTEXT_TOKEN_LIMIT = 32_768
DEFAULT_MAX_COMPLETION_TOKENS = 4096
_POLICY_CONNECTION_LIMIT = 2048
_POLICY_CONNECTION_LIMIT = 100_000
_POLICY_MAX_RETRIES = 1
_POLICY_HTTP_TIMEOUT = httpx.Timeout(connect=30, read=10 * 60, write=30, pool=30)

Expand Down
5 changes: 4 additions & 1 deletion src/art/trajectories/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -1223,7 +1223,10 @@ def _intern_source_graph(self) -> TokenizedTrajectoryGroup[TokenizedTrajectoryT]
for tokenized, trajectory in zip(
self.trajectories, self.trajectory_group.trajectories, strict=True
):
if tokenized.trajectory.model_dump() != trajectory.model_dump():
if (
tokenized.trajectory is not trajectory
and tokenized.trajectory.model_dump() != trajectory.model_dump()
):
raise ValueError("Tokenized trajectory does not match its source group")
tokenized.trajectory = trajectory
if isinstance(tokenized, TokenizedTrajectory):
Expand Down
4 changes: 1 addition & 3 deletions src/art/trajectories/_history.py
Original file line number Diff line number Diff line change
Expand Up @@ -392,9 +392,7 @@ def _chat_generation_tokens(
)
if choice is None:
raise ValueError("Chat choice source index is out of bounds")
prompt, output, _ = _chat_choice_tokens(
choice, exchange.response.model_dump(mode="python")
)
prompt, output, _ = _chat_choice_tokens(choice, exchange.response)
cache[key] = prompt, output
return cache[key]

Expand Down
40 changes: 21 additions & 19 deletions src/art/trajectories/_serialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,9 @@ def _intern_strings(value: object, pool: _StringPool | None = None) -> None:
def _intern_value(value: object, pool: _StringPool, memo: dict[int, object]) -> object:
if isinstance(value, str):
return pool.setdefault(value, value)
if isinstance(value, (bytes, bytearray, memoryview)) or value is None:
if value is None or isinstance(
value, (bytes, bytearray, memoryview, bool, int, float, complex)
):
return value

value_id = id(value)
Expand All @@ -59,15 +61,6 @@ def _intern_value(value: object, pool: _StringPool, memo: dict[int, object]) ->
if isinstance(value, _StringInterningModel):
value._mark_pickle_strings_interned()
return value
if is_dataclass(value) and type(value).__module__.startswith("art.trajectories"):
memo[value_id] = value
for field in fields(value):
object.__setattr__(
value,
field.name,
_intern_value(getattr(value, field.name), pool, memo),
)
return value
if isinstance(value, dict):
memo[value_id] = value
_intern_mapping(cast(dict[object, object], value), pool, memo)
Expand Down Expand Up @@ -95,21 +88,30 @@ def _intern_value(value: object, pool: _StringPool, memo: dict[int, object]) ->
result = frozenset(_intern_value(item, pool, memo) for item in value)
memo[value_id] = result
return result
if is_dataclass(value) and type(value).__module__.startswith("art.trajectories"):
memo[value_id] = value
for field in fields(value):
object.__setattr__(
value,
field.name,
_intern_value(getattr(value, field.name), pool, memo),
)
return value
return value


def _intern_mapping(
value: dict[object, object], pool: _StringPool, memo: dict[int, object]
) -> None:
items = [
(
_intern_value(key, pool, memo) if isinstance(key, str) else key,
_intern_value(item, pool, memo),
)
for key, item in value.items()
]
value.clear()
value.update(items)
replacements: list[tuple[str, str]] = []
for key, item in value.items():
if isinstance(key, str):
interned = pool.setdefault(key, key)
if interned is not key:
replacements.append((key, interned))
value[key] = _intern_value(item, pool, memo)
for key, interned in replacements:
value[interned] = value.pop(key)


def serialize_messages_and_choices(items: list[Any]) -> list[dict[str, Any]]:
Expand Down
19 changes: 8 additions & 11 deletions src/art/trajectories/_tokenize.py
Original file line number Diff line number Diff line change
Expand Up @@ -454,7 +454,7 @@ def _logprob_values(values: object) -> list[float]:
return []
result: list[float] = []
for value in values:
logprob = _dump(value).get("logprob")
logprob = _field(value, "logprob")
if not isinstance(logprob, (int, float)) or isinstance(logprob, bool):
return []
result.append(float(logprob))
Expand All @@ -473,17 +473,15 @@ def _chat_logprob_entries(choice: Choice) -> list[object]:
def _chat_choice_output_tokens(
choice: Choice,
) -> tuple[list[int] | None, list[float]]:
choice_data = _dump(choice)
token_ids = _exact_token_ids(
choice_data.get("token_ids"),
_field(choice, "token_ids"),
field="Chat Completions token_ids",
)
values = _chat_logprob_entries(choice)
message = _dump(choice.message)
if token_ids == [] and (
values
or any(
message.get(key)
_field(choice.message, key)
for key in (
"content",
"refusal",
Expand All @@ -507,12 +505,11 @@ def _chat_choice_output_tokens(


def _chat_choice_tokens(
choice: Choice, response_data: dict[str, Any]
choice: Choice, response: object
) -> tuple[list[int] | None, list[int] | None, list[float]]:
choice_data = _dump(choice)
prompt = choice_data.get("prompt_token_ids")
prompt = _field(choice, "prompt_token_ids")
if prompt is None:
prompt = response_data.get("prompt_token_ids")
prompt = _field(response, "prompt_token_ids")
prompt_ids = _exact_token_ids(
prompt,
field="Chat Completions prompt_token_ids",
Expand All @@ -531,7 +528,7 @@ def _chat_tokens(
) -> tuple[list[int] | None, list[int] | None, list[float]]:
if len(response.choices) != 1:
raise ValueError("Trajectory tokenization requires exactly one response choice")
return _chat_choice_tokens(response.choices[0], _dump(response))
return _chat_choice_tokens(response.choices[0], response)


def _completion_evidence(
Expand Down Expand Up @@ -3028,7 +3025,7 @@ def _chat_source_prompt_tokens(source: object) -> list[int] | None:
choice = next(
item for item in exchange.response.choices if item.index == choice_index
)
prompt, _, _ = _chat_choice_tokens(choice, _dump(exchange.response))
prompt, _, _ = _chat_choice_tokens(choice, exchange.response)
return prompt
if isinstance(exchange, ResponsesExchange):
generation_index = getattr(source, "generation_index", None)
Expand Down
5 changes: 4 additions & 1 deletion src/art/trajectories/tensors.py
Original file line number Diff line number Diff line change
Expand Up @@ -237,7 +237,10 @@ def bind_source_group(self) -> Self:
for tensorized, trajectory in zip(
self.trajectories, self.trajectory_group.trajectories, strict=True
):
if tensorized.trajectory.model_dump() != trajectory.model_dump():
if (
tensorized.trajectory is not trajectory
and tensorized.trajectory.model_dump() != trajectory.model_dump()
):
raise ValueError(
"Tensorized trajectory does not match its source group"
)
Expand Down
8 changes: 4 additions & 4 deletions tests/unit/test_tau_bench_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,8 +42,8 @@ def __init__(self, **kwargs: Any) -> None:

limits = seen["transport_kwargs"]["limits"]
assert isinstance(limits, httpx.Limits)
assert limits.max_connections == 512
assert limits.max_keepalive_connections == 512
assert limits.max_connections == 100_000
assert limits.max_keepalive_connections == 100_000
assert seen["transport_kwargs"]["retries"] == 2
assert isinstance(seen["timeout"], httpx.Timeout)

Expand Down Expand Up @@ -329,8 +329,8 @@ async def test_rollout_supports_string_model_args(
write=30,
pool=30,
)
assert http_client.limits.max_connections == 2048
assert http_client.limits.max_keepalive_connections == 2048
assert http_client.limits.max_connections == 100_000
assert http_client.limits.max_keepalive_connections == 100_000


@pytest.mark.asyncio
Expand Down
14 changes: 14 additions & 0 deletions tests/unit/trajectories/test_history.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,20 @@ def _completion(
)


def test_chat_history_extracts_token_evidence_without_dumping_full_responses(
monkeypatch: pytest.MonkeyPatch,
) -> None:
trajectory = _growing_chat_trajectory(3, exact_tokens=True)

def unexpected_dump(*_: object, **__: object) -> object:
raise AssertionError("chat token evidence must not dump the full response")

monkeypatch.setattr(ChatCompletion, "model_dump", unexpected_dump)

histories = trajectory.histories()
assert len(histories) == 1


def _message() -> MessagesExchange:
start, end = _times()
return MessagesExchange(
Expand Down
21 changes: 21 additions & 0 deletions tests/unit/trajectories/test_tensorized_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -302,6 +302,27 @@ def test_tensorized_pickle_retains_sources_without_tokenized_intermediate() -> N
assert torch.equal(restored.tokens, tensorized.tokens)


def test_tensorized_group_skips_equality_dump_for_identical_source(
monkeypatch: pytest.MonkeyPatch,
) -> None:
source = art.Trajectory()
tokenized = tr.TokenizedTrajectory(
**_tokenized_history().model_dump(), trajectory=source
)
tensorized = tokenized.tensorize()

def unexpected_dump(*_: object, **__: object) -> object:
raise AssertionError("identical source trajectories must not be dumped")

monkeypatch.setattr(art.Trajectory, "model_dump", unexpected_dump)
group = tr.TensorizedTrajectoryGroup[tr.TensorizedTrajectory](
trajectory_group=art.TrajectoryGroup([source]),
trajectories=[tensorized],
)

assert group.trajectories[0].trajectory is source


def test_tensorized_group_pydantic_and_cloudpickle_round_trips() -> None:
cloudpickle = pytest.importorskip("cloudpickle")
trajectory = _trajectory()
Expand Down
20 changes: 20 additions & 0 deletions tests/unit/trajectories/test_tokenized_models.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
import math
import pickle

import pytest

import art
import art.trajectories as tr

Expand Down Expand Up @@ -72,6 +74,24 @@ def test_nested_tokenized_models_nan_json_round_trip() -> None:
assert restored.metadata == group.metadata


def test_tokenized_group_skips_equality_dump_for_identical_source(
monkeypatch: pytest.MonkeyPatch,
) -> None:
source = art.Trajectory()
tokenized = tr.TokenizedTrajectory(**_history().model_dump(), trajectory=source)

def unexpected_dump(*_: object, **__: object) -> object:
raise AssertionError("identical source trajectories must not be dumped")

monkeypatch.setattr(art.Trajectory, "model_dump", unexpected_dump)
group = tr.TokenizedTrajectoryGroup[tr.TokenizedTrajectory](
trajectory_group=art.TrajectoryGroup([source]),
trajectories=[tokenized],
)

assert group.trajectories[0].trajectory is source


def test_public_group_tokenization_nan_json_round_trip() -> None:
from datetime import datetime

Expand Down
Loading