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:
@@ -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"]}
|
||||||
|
|||||||
Reference in New Issue
Block a user