File size: 2,619 Bytes
25f9bfc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78

# worker/rs2s_tasks/runtime_cli.py
import os
import torch
from lightning.pytorch.cli import LightningCLI
from lightning.pytorch import Trainer
from pathlib import Path
from transformers_model.model import LitVanillaTransformer
from transformers_model.smiles_datamodule import LitSmilesDataset

FORWARD_CKPT_PATH = os.getenv(
    "FORWARD_CKPT_PATH",
    "rs2s_tasks/epoch=59-val_accuracy=0.6476.ckpt"
)

REQUESTED_DEVICE = os.getenv("DEVICE", "").lower()
def _choose_accelerator():
    if REQUESTED_DEVICE == "cuda" and torch.cuda.is_available():
        return "gpu"
    return "cpu"

_CLI = None
_TRAINER = None

def get_cli_runtime():
    """Instantiate LightningCLI once and cache model/datamodule/trainer."""
    global _CLI, _TRAINER
    if _CLI is None:
        # Build CLI programmatically (no argv parsing). We mimic your predict setup.
        _CLI = LightningCLI(
            model_class=LitVanillaTransformer,
            datamodule_class=LitSmilesDataset,
            subclass_mode_model=False,
            subclass_mode_data=False,
            run=False,  # <--- do not launch fit/test/predict automatically
            args=[
                f"--model.vocab_path={str(Path(__file__).with_name('vocab.txt'))}",
                "--model.task=forward",
                "--model.device=cpu",
                f"--data.vocab_path={str(Path(__file__).with_name('vocab.txt'))}",
                "--data.batch_size=512",
                "--seed_everything=42",
            ],
        )

        # Load weights exactly like CLI predict
        _CLI.model = LitVanillaTransformer.load_from_checkpoint(
            checkpoint_path=FORWARD_CKPT_PATH,
            strict=False,
            map_location="cpu",  # safe default
            vocab_path=str(Path(__file__).with_name("vocab.txt")),
            task="forward",
            device="cpu",
        )
        _CLI.model.eval().freeze()

        _CLI.datamodule = LitSmilesDataset(
            batch_size=512,
            vocab_path=str(Path(__file__).with_name("vocab.txt")),
        )

        # Build a lightweight Trainer for predict calls
        _TRAINER = Trainer(
            accelerator="cpu",  #_choose_accelerator(),
            devices=1,
            logger=False,
            enable_checkpointing=False,
            inference_mode=True,
        )

        # Place model on the selected device (after construction)
        #device = torch.device("cuda" if _TRAINER.accelerator.strategy.root_device.type == "cuda" else "cpu")
        device = "cpu"
        _CLI.model.to(device).eval()
        torch.set_grad_enabled(False)

    return _CLI, _TRAINER