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
1 change: 1 addition & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ MODEL_NAME=google/gemini-2.5-flash
# LAYA_DEVICE=cpu # cuda when available
# LAYA_MAX_LEN=1024
# LAYA_HEAD_MAX_LEN=512 # raise for choice questions with many options
# LAYA_MPS_AMP_MIN_ROWS=5 # on MPS, fp16 from this many questions (laya's default); unset keeps fp32

# ---- Cua-S1 Nano (the in-process option scorer behind --model cua; needs `uv sync --extra cua`) ----
# CUA_S1_CHECKPOINT=cua-ai/cua-s1-nano-0.1 # a Hugging Face id, or a local directory holding <subfolder>/
Expand Down
7 changes: 5 additions & 2 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -20,8 +20,11 @@ The format follows [Keep a Changelog](https://keepachangelog.com/en/1.1.0/); ver

### Changed

- `--model laya` loads in about 3 s instead of about 35 s: the encoder is built with transformers' weight init
off, since the checkpoint replaces every weight. Weights and answers are unchanged.
- `--model laya` no longer draws the encoder's random weights before the checkpoint replaces them, which took most
of a load of about 40 s on CPU. The `laya` extra now needs laya 0.3.9 or later, which skips the draw itself,
and the lock moves from 0.3.5 to 0.3.20. Weights and answers are unchanged: laya 0.3.10 and later run a request
of five or more questions in fp16 on MPS, which moves the answers, so `--model laya` keeps such requests in fp32
unless `LAYA_MPS_AMP_MIN_ROWS` is set.
- `--model` picks the model on every agent, on `decide` and on `probe`: `jev`, `laya`, `cua`, `llm`, `random` or
`rule`. The results table's column, the replay page's badge data and a browser run's `answer.json` name it
`model` as well; the replay still reads the `slot` key of records written by 0.1.0.
Expand Down
1 change: 1 addition & 0 deletions docs/configuration.md
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@ Variables can be exported in your shell or placed in a `.env` file at the root o
| `LAYA_DEVICE` | `laya` model | `(library default)` | PyTorch device for Laya model evaluation; passes None so the library selects CUDA, MPS, or CPU. |
| `LAYA_MAX_LEN` | `laya` model | `(checkpoint default)` | Maximum token sequence length for Laya state representation; overrides checkpoint window only when set. |
| `LAYA_HEAD_MAX_LEN` | `laya` model | `(checkpoint default)` | Maximum token sequence length for Laya decision head options; overrides checkpoint window only when set. |
| `LAYA_MPS_AMP_MIN_ROWS` | `laya` model | *(unset: fp32)* | Laya's own variable: on MPS, requests with at least this many questions run in fp16. Unset, `--model laya` keeps every request in fp32, as on CPU; `5` is Laya's default. fp16 moves the probabilities and can flip a close decision. |
| `CUA_S1_CHECKPOINT` | `cua` model | `cua-ai/cua-s1-nano-0.1` | Hugging Face checkpoint ID or local directory for Cua-S1 Nano option scorer. |
| `CUA_S1_SUBFOLDER` | `cua` model | `text` | Subfolder within checkpoint directory containing text option scoring weights. |
| `CUA_S1_DEVICE` | `cua` model | `auto` | PyTorch device used for Cua-S1 Nano evaluation (`auto`, `cpu`, `cuda`, or `mps`). |
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,7 @@ alfworld-visual = [
"torchvision>=0.15",
]
report = ["pillow>=10", "playwright>=1.45"] # evals.replay: pages, GIFs; never imported by a runner
laya = ["laya>=0.3.4"] # the in-process decision model behind --model laya; pulls torch and transformers
laya = ["laya>=0.3.9"] # the in-process decision model behind --model laya; pulls torch and transformers
cua = [ # Cua-S1 Nano behind --model cua; pinned to the Cua PR that ships the checkpoints (trycua/cua#4023), pulls torch
"cua-s1 @ git+https://github.com/trycua/cua.git@aea61b6eb97e2d8c0f6f71eb804e5769fe910af4#subdirectory=libs/cua-s1/python",
"huggingface-hub>=0.24",
Expand Down
26 changes: 8 additions & 18 deletions s1a/decision_models/laya.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,8 @@
from __future__ import annotations

import asyncio
import importlib
import os
import time
from contextlib import AbstractContextManager, nullcontext
from importlib import metadata
from typing import Any

Expand All @@ -23,6 +21,7 @@

LAYA_DEFAULT_MODEL = "convaiinnovations/laya"
LAYA_DEFAULT_MAX_LEN = 512 # the window Laya assumes when a checkpoint config names none
LAYA_MPS_FP32_ROWS = 10**9 # no request has this many questions, so Laya never switches to fp16 on MPS


def laya_question(question: Question) -> Json:
Expand All @@ -34,19 +33,6 @@ def laya_question(question: Question) -> Json:
return {"type": "noul", "instructions": question.question, **criteria}


def without_weight_init() -> AbstractContextManager[Any]:
"""transformers' ``no_init_weights``. ``laya.load`` builds the encoder from its config, which draws every weight
at random (about 30 s of the load on CPU), then loads the checkpoint over all of them with ``strict=True``, so
the draw is thrown away. The helper sits in ``transformers.initialization`` from 5.0 and in
``transformers.modeling_utils`` before; without either the load runs as it is."""
for module in ("transformers.initialization", "transformers.modeling_utils"):
try:
return importlib.import_module(module).no_init_weights()
except (ImportError, AttributeError):
continue
return nullcontext()


class LayaModel(DecisionModel):
"""Laya's ``Agent`` (or anything with ``system_one(state, questions)`` and a ``cfg``) behind the interface."""

Expand Down Expand Up @@ -103,7 +89,8 @@ def _check_the_window(self, usage: Usage, questions: int) -> None:
@classmethod
def from_env(cls) -> "LayaModel":
"""``LAYA_MODEL`` (a hub id or a path), ``LAYA_SUBFOLDER``, ``LAYA_DEVICE``; ``LAYA_MAX_LEN`` and
``LAYA_HEAD_MAX_LEN`` override the checkpoint's window."""
``LAYA_HEAD_MAX_LEN`` override the checkpoint's window. On MPS the model answers in fp32 whatever the
number of questions, unless ``LAYA_MPS_AMP_MIN_ROWS`` (Laya's own variable) is set."""
try:
import laya
except ImportError as exc:
Expand All @@ -113,8 +100,7 @@ def from_env(cls) -> "LayaModel":
) from exc
model = os.getenv("LAYA_MODEL") or LAYA_DEFAULT_MODEL
subfolder = os.getenv("LAYA_SUBFOLDER") or None
with without_weight_init():
agent = laya.load(model, device=os.getenv("LAYA_DEVICE") or None, subfolder=subfolder)
agent = laya.load(model, device=os.getenv("LAYA_DEVICE") or None, subfolder=subfolder)
if not callable(getattr(agent, "system_one", None)):
try:
version = metadata.version("laya")
Expand All @@ -127,6 +113,10 @@ def from_env(cls) -> "LayaModel":
"the s1a laya model needs that method"
),
)
# From 0.3.10 Laya runs a request of five or more questions in fp16 on MPS. That moves the answers, enough
# to flip a close decision, so one browser episode would mix both precisions.
if not os.getenv("LAYA_MPS_AMP_MIN_ROWS") and hasattr(agent, "mps_amp_min_rows"):
agent.mps_amp_min_rows = LAYA_MPS_FP32_ROWS
for key, variable in (("max_len", "LAYA_MAX_LEN"), ("head_max_len", "LAYA_HEAD_MAX_LEN")):
value = os.getenv(variable)
if value:
Expand Down
5 changes: 1 addition & 4 deletions tests/test_decision_models_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@

import os
import sys
from contextlib import nullcontext
from types import SimpleNamespace
from unittest import TestCase
from unittest.mock import patch
Expand All @@ -32,9 +31,7 @@ def test_every_name_builds_its_class(self) -> None:
fake_laya = SimpleNamespace(
load=lambda *a, **k: SimpleNamespace(cfg={}, system_one=lambda state, questions: {})
)
# transformers' helper stubbed so torch is not imported inside patch.dict: see TestFromEnv in the laya tests.
modules = {"laya": fake_laya, "transformers.initialization": SimpleNamespace(no_init_weights=nullcontext)}
with patch.dict(sys.modules, modules), patch.dict(os.environ, {"LAYA_SUBFOLDER": ""}):
with patch.dict(sys.modules, {"laya": fake_laya}), patch.dict(os.environ, {"LAYA_SUBFOLDER": ""}):
self.assertIsInstance(build_model("laya"), LayaModel)
self.assertIsInstance(build_model("random", seed=3), RandomModel)
rule = build_model("rule", rule=("always-inc", lambda state, options: "inc"))
Expand Down
92 changes: 37 additions & 55 deletions tests/test_decision_models_laya.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,12 @@
# coding: utf-8
"""``LayaModel`` over a fake ``laya.Agent`` (no torch): the contract, the question mapping, the error wrap,
the filled-window error, and ``from_env`` with and without the extra and with the weight init off."""
the filled-window error, and ``from_env`` with and without the extra."""

from __future__ import annotations

import os
import sys
import time
from collections.abc import Iterator
from contextlib import contextmanager, nullcontext
from types import SimpleNamespace
from typing import Any
from unittest import IsolatedAsyncioTestCase, TestCase
Expand Down Expand Up @@ -163,14 +161,6 @@ async def test_the_window_scales_with_the_number_of_questions(self) -> None:


class TestFromEnv(TestCase):
def setUp(self) -> None:
# A stand-in for transformers' helper, so no test imports torch: patch.dict drops a torch imported inside it
# from sys.modules, and importing torch a second time in one process crashes it.
helper = {"transformers.initialization": SimpleNamespace(no_init_weights=nullcontext)}
stub = patch.dict(sys.modules, helper)
stub.start()
self.addCleanup(stub.stop)

def test_without_the_extra_it_is_a_config_error_naming_the_extra(self) -> None:
with patch.dict(sys.modules, {"laya": None}):
with self.assertRaises(BaseError) as caught:
Expand All @@ -180,13 +170,48 @@ def test_without_the_extra_it_is_a_config_error_naming_the_extra(self) -> None:

def test_an_agent_without_system_one_is_a_config_error_naming_the_method(self) -> None:
fake_laya = SimpleNamespace(load=lambda *a, **k: SimpleNamespace(cfg={}))
with patch.dict(sys.modules, {"laya": fake_laya}), patch.dict(os.environ, {"LAYA_SUBFOLDER": ""}):
# The installed version is pinned so the message reads the same with and without the extra.
with (
patch.dict(sys.modules, {"laya": fake_laya}),
patch.dict(os.environ, {"LAYA_SUBFOLDER": ""}),
patch.object(laya_module.metadata, "version", return_value="0.3.0"),
):
with self.assertRaises(BaseError) as caught:
LayaModel.from_env()
self.assertEqual(caught.exception.status, StatusCode.MODEL_SERVICE_CONFIG_ERROR)
self.assertIn("system_one", str(caught.exception))
self.assertIn("laya 0.3.0", str(caught.exception))

def test_a_laya_without_package_metadata_is_named_unknown(self) -> None:
fake_laya = SimpleNamespace(load=lambda *a, **k: SimpleNamespace(cfg={}))
with (
patch.dict(sys.modules, {"laya": fake_laya}),
patch.dict(os.environ, {"LAYA_SUBFOLDER": ""}),
patch.object(laya_module.metadata, "version", side_effect=laya_module.metadata.PackageNotFoundError),
):
with self.assertRaises(BaseError) as caught:
LayaModel.from_env()
self.assertIn("laya unknown", str(caught.exception))

def test_mps_stays_in_fp32_unless_the_laya_variable_is_set(self) -> None:
def from_env(env: dict[str, str]) -> Any:
agent = FakeLayaAgent()
agent.mps_amp_min_rows = 5 # what laya sets from 0.3.10: fp16 from five questions on MPS
fake_laya = SimpleNamespace(load=lambda *a, **k: agent)
env = {"LAYA_SUBFOLDER": "", "LAYA_MPS_AMP_MIN_ROWS": "", **env}
with patch.dict(sys.modules, {"laya": fake_laya}), patch.dict(os.environ, env):
return LayaModel.from_env()._agent

self.assertEqual(from_env({}).mps_amp_min_rows, laya_module.LAYA_MPS_FP32_ROWS)
self.assertEqual(from_env({"LAYA_MPS_AMP_MIN_ROWS": "5"}).mps_amp_min_rows, 5)

def test_a_laya_before_the_mps_gate_gets_no_such_attribute(self) -> None:
agent = FakeLayaAgent() # laya 0.3.9 has no mps_amp_min_rows
with patch.dict(sys.modules, {"laya": SimpleNamespace(load=lambda *a, **k: agent)}):
with patch.dict(os.environ, {"LAYA_SUBFOLDER": "", "LAYA_MPS_AMP_MIN_ROWS": ""}):
LayaModel.from_env()
self.assertFalse(hasattr(agent, "mps_amp_min_rows"))

def test_the_env_names_the_checkpoint_and_overrides_the_window(self) -> None:
loads: list[tuple[Any, ...]] = []

Expand All @@ -207,49 +232,6 @@ def load(model: str, device: Any = None, token: Any = None, subfolder: Any = Non
self.assertEqual(decision_model.model, "convaiinnovations/laya/multilingual")
self.assertEqual(decision_model._agent.cfg, {"max_len": 1024, "head_max_len": 512})

def test_the_checkpoint_loads_with_the_weight_init_off(self) -> None:
events: list[str] = []

@contextmanager
def no_init_weights() -> Iterator[None]:
events.append("off")
yield
events.append("on")

def load(*args: Any, **kwargs: Any) -> FakeLayaAgent:
events.append("load")
return FakeLayaAgent()

modules = {
"laya": SimpleNamespace(load=load),
"transformers.initialization": SimpleNamespace(no_init_weights=no_init_weights),
}
with patch.dict(sys.modules, modules), patch.dict(os.environ, {"LAYA_SUBFOLDER": ""}):
LayaModel.from_env()
self.assertEqual(events, ["off", "load", "on"])

def test_the_4x_location_of_the_helper_is_used_when_the_5x_one_is_missing(self) -> None:
@contextmanager
def no_init_weights() -> Iterator[str]:
yield "4.x"

modules = {
"transformers.initialization": None,
"transformers.modeling_utils": SimpleNamespace(no_init_weights=no_init_weights),
}
with patch.dict(sys.modules, modules), laya_module.without_weight_init() as entered:
self.assertEqual(entered, "4.x")

def test_without_the_helper_the_checkpoint_still_loads(self) -> None:
modules = {
"laya": SimpleNamespace(load=lambda *a, **k: FakeLayaAgent()),
"transformers.initialization": None,
"transformers.modeling_utils": SimpleNamespace(),
}
with patch.dict(sys.modules, modules), patch.dict(os.environ, {"LAYA_SUBFOLDER": ""}):
decision_model = LayaModel.from_env()
self.assertIsInstance(decision_model._agent, FakeLayaAgent)

def test_the_defaults_when_the_env_is_empty(self) -> None:
env = {"LAYA_MODEL": "", "LAYA_SUBFOLDER": "", "LAYA_DEVICE": "", "LAYA_MAX_LEN": "", "LAYA_HEAD_MAX_LEN": ""}
with patch.dict(sys.modules, {"laya": SimpleNamespace(load=lambda *a, **k: FakeLayaAgent())}):
Expand Down
8 changes: 4 additions & 4 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading