From 86b5c3389785b150dee30929a3aaed942e4acbd5 Mon Sep 17 00:00:00 2001 From: cacheline999 <326908201+cacheline999@users.noreply.github.com> Date: Sun, 27 Sep 2026 13:28:42 +0800 Subject: [PATCH] [Fix] Load Laya without the throwaway random weight init laya.load builds the ModernBERT encoder from its config, so transformers draws every weight at random before the checkpoint is loaded over all of them with strict=True. That draw is ~33 of the ~35 s load. Run the load under transformers' no_init_weights (transformers.initialization in 5.x, transformers.modeling_utils in 4.x, plain load if neither exists). The from_env tests stub the helper: importing torch inside patch.dict(sys.modules) and importing it again later segfaults. --- CHANGELOG.md | 2 + s1a/decision_models/laya.py | 18 ++++++++- tests/test_decision_models_factory.py | 5 ++- tests/test_decision_models_laya.py | 55 ++++++++++++++++++++++++++- 4 files changed, 77 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5520df2..7b1eca6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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. diff --git a/s1a/decision_models/laya.py b/s1a/decision_models/laya.py index 1fb786d..b4178a2 100644 --- a/s1a/decision_models/laya.py +++ b/s1a/decision_models/laya.py @@ -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 @@ -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.""" @@ -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") diff --git a/tests/test_decision_models_factory.py b/tests/test_decision_models_factory.py index d38866e..3c85f6d 100644 --- a/tests/test_decision_models_factory.py +++ b/tests/test_decision_models_factory.py @@ -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 @@ -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")) diff --git a/tests/test_decision_models_laya.py b/tests/test_decision_models_laya.py index 2740558..992924d 100644 --- a/tests/test_decision_models_laya.py +++ b/tests/test_decision_models_laya.py @@ -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 @@ -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: @@ -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())}):