fix(intern-decision-serve): register the checkpoint's inference module before executing it

Its dataclasses use postponed annotations and look their module up in sys.modules while the
class is built; importing by file path without registering it failed startup (closed).
This commit is contained in:
vh
2026-09-30 09:08:19 -07:00
parent cc211ef10c
commit f7415db5c9
2 changed files with 18 additions and 5 deletions
@@ -13,6 +13,7 @@ import hashlib
import importlib.metadata import importlib.metadata
import importlib.util import importlib.util
import logging import logging
import sys
import traceback import traceback
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
@@ -46,11 +47,21 @@ def _sha256(path: Path) -> str:
return hashlib.sha256(path.read_bytes()).hexdigest() return hashlib.sha256(path.read_bytes()).hexdigest()
INFERENCE_MODULE = "intern_decision_inference"
def _import_inference(path: Path): def _import_inference(path: Path):
"""The checkpoint's own inference.py, imported by file path under a private module name.""" """The checkpoint's own inference.py, imported by file path under a private module name. It is
spec = importlib.util.spec_from_file_location("intern_decision_inference", path) registered in sys.modules BEFORE it runs: its dataclasses (with `from __future__ import
annotations`) look their module up there while the class is being built."""
spec = importlib.util.spec_from_file_location(INFERENCE_MODULE, path)
module = importlib.util.module_from_spec(spec) module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module) sys.modules[INFERENCE_MODULE] = module
try:
spec.loader.exec_module(module)
except BaseException:
sys.modules.pop(INFERENCE_MODULE, None)
raise
return module return module
@@ -89,15 +89,17 @@ def reset_fake_torch():
# A fake checkpoint: snapshots/<REVISION>/inference.py defining a DecisionEngine shaped like the real one. # A fake checkpoint: snapshots/<REVISION>/inference.py defining a DecisionEngine shaped like the real one.
# --------------------------------------------------------------------------------------------- # ---------------------------------------------------------------------------------------------
FAKE_INFERENCE = textwrap.dedent(''' FAKE_INFERENCE = textwrap.dedent('''
from __future__ import annotations # as in the real one: dataclasses then look the module up in sys.modules
import json import json
from dataclasses import dataclass
MODEL_NAME = "Intern-Decision-4B" MODEL_NAME = "Intern-Decision-4B"
EVENTS = [] EVENTS = []
WARMUP_SHIFT = {"after_swap": 0.0} WARMUP_SHIFT = {"after_swap": 0.0}
TOKEN_SKEW = {"n": 0} TOKEN_SKEW = {"n": 0}
@dataclass(frozen=True) # the real inference.py defines one: needs sys.modules at exec time
class Compiled: class Compiled:
def __init__(self, messages): messages: list
self.messages = messages
def validate_request(request): def validate_request(request):
return {"state": request["state"], "questions": request["questions"]} return {"state": request["state"], "questions": request["questions"]}