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
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@ 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` 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
18 changes: 17 additions & 1 deletion s1a/decision_models/laya.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,10 @@
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 @@ -32,6 +34,19 @@ 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 @@ -98,7 +113,8 @@ def from_env(cls) -> "LayaModel":
) from exc
model = os.getenv("LAYA_MODEL") or LAYA_DEFAULT_MODEL
subfolder = os.getenv("LAYA_SUBFOLDER") or None
agent = laya.load(model, device=os.getenv("LAYA_DEVICE") or None, subfolder=subfolder)
with without_weight_init():
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 Down
5 changes: 4 additions & 1 deletion tests/test_decision_models_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

import os
import sys
from contextlib import nullcontext
from types import SimpleNamespace
from unittest import TestCase
from unittest.mock import patch
Expand All @@ -31,7 +32,9 @@ def test_every_name_builds_its_class(self) -> None:
fake_laya = SimpleNamespace(
load=lambda *a, **k: SimpleNamespace(cfg={}, system_one=lambda state, questions: {})
)
with patch.dict(sys.modules, {"laya": fake_laya}), patch.dict(os.environ, {"LAYA_SUBFOLDER": ""}):
# 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": ""}):
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
55 changes: 54 additions & 1 deletion tests/test_decision_models_laya.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,14 @@
# 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."""
the filled-window error, and ``from_env`` with and without the extra and with the weight init off."""

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 @@ -161,6 +163,14 @@ 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 Down Expand Up @@ -197,6 +207,49 @@ 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