File size: 5,592 Bytes
bf314e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
#!/usr/bin/env python3
import argparse
import pickle
from pathlib import Path
import torch
from omegaconf import OmegaConf, DictConfig, ListConfig

ROOT = Path(__file__).resolve().parents[2]
BACKBONE_MOE = "model.uma_escn_moe"
BACKBONE_MD = "model.uma_escn_md"
BACKBONE_BASE = "model.base"

RULES = [
    ("onescience.models.UMA.models.base", BACKBONE_BASE),
    ("onescience.models.UMA.base", BACKBONE_BASE),
    ("onescience.models.UMA.models.uma.escn_md", BACKBONE_MD),
    ("onescience.models.UMA.models.uma.escn_moe", BACKBONE_MOE),
    ("onescience.models.UMA.uma.escn_md", BACKBONE_MD),
    ("onescience.models.UMA.uma.escn_moe", BACKBONE_MOE),
    ("onescience.models.UMA.uma_escn_md", BACKBONE_MD),
    ("onescience.models.UMA.uma_escn_moe", BACKBONE_MOE),
    ("onescience.utils.uma.models.uma.escn_md", BACKBONE_MD),
    ("onescience.utils.uma.models.uma.escn_moe", BACKBONE_MOE),
    ("fairchem.core.models.base", BACKBONE_BASE),
    ("fairchem.core.models.uma.escn_md", BACKBONE_MD),
    ("fairchem.core.models.uma.escn_moe", BACKBONE_MOE),
    ("onescience.models.UMA.units", "onescience.utils.uma.units"),
    ("onescience.models.UMA.modules.head", "onescience.modules.head.uma_head"),
    ("onescience.models.UMA.modules.loss", "onescience.modules.loss.uma_loss"),
    ("onescience.models.UMA.common", "onescience.utils.uma.common"),
    ("onescience.utils.uma.modules.head", "onescience.modules.head.uma_head"),
    ("onescience.utils.uma.modules.loss", "onescience.modules.loss.uma_loss"),
    ("fairchem.core.units", "onescience.utils.uma.units"),
    ("fairchem.core.modules.head", "onescience.modules.head.uma_head"),
    ("fairchem.core.modules.loss", "onescience.modules.loss.uma_loss"),
    ("fairchem.core.components", "onescience.utils.uma.components"),
    ("fairchem.core.common", "onescience.utils.uma.common"),
    ("fairchem.core", "onescience.utils.uma"),
    # 兜底:修复之前被错误转换的中间状态
    ("onescience.modules.normalization", "onescience.utils.uma.normalization"),
]

LEGACY_PREFIX = (
    "onescience.models.UMA.models.",
    "onescience.models.UMA.uma.",
    "onescience.utils.uma.models.",
    "onescience.utils.uma.modules.head.",
    "onescience.utils.uma.modules.loss.",
    "fairchem.core.models.",
    "fairchem.core.modules.head.",
    "fairchem.core.modules.loss.",
    "onescience.modules.normalization.",
)

def remap(s: str) -> str:
    prev = None
    cur = s
    while cur != prev:
        prev = cur
        for old, new in RULES:
            if cur == old or cur.startswith(old + "."):
                cur = new + cur[len(old):]
                break
        # 兜底:修复之前错误转换出的中间状态
        if cur.startswith("onescience.modules.loss.") and not cur.startswith("onescience.modules.loss.uma_loss."):
            cur = "onescience.modules.loss.uma_loss." + cur[len("onescience.modules.loss."):]
        if cur.startswith("onescience.modules.head.") and not cur.startswith("onescience.modules.head.uma_head."):
            cur = "onescience.modules.head.uma_head." + cur[len("onescience.modules.head."):]
    return cur

def to_py(x):
    if isinstance(x, (DictConfig, ListConfig)):
        return OmegaConf.to_container(x, resolve=False)
    return x

def walk(x):
    if isinstance(x, str):
        return remap(x)
    if isinstance(x, dict):
        return {k: walk(v) for k, v in x.items()}
    if isinstance(x, list):
        return [walk(v) for v in x]
    return x

def find_legacy(x, path="root"):
    out = []
    if isinstance(x, str):
        if x.startswith(LEGACY_PREFIX):
            out.append((path, x))
    elif isinstance(x, dict):
        for k, v in x.items():
            out.extend(find_legacy(v, f"{path}.{k}"))
    elif isinstance(x, list):
        for i, v in enumerate(x):
            out.extend(find_legacy(v, f"{path}[{i}]"))
    return out

class CustomUnpickler(pickle.Unpickler):
    def find_class(self, module_name, class_name):
        mapped = remap(module_name)
        try:
            return super().find_class(mapped, class_name)
        except ModuleNotFoundError:
            return super().find_class(module_name, class_name)

class CustomPickleModule:
    Unpickler = CustomUnpickler

def gf(obj, k):
    return obj[k] if isinstance(obj, dict) else getattr(obj, k)

def sf(obj, k, v):
    if isinstance(obj, dict):
        obj[k] = v
    else:
        setattr(obj, k, v)

def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("src")
    ap.add_argument("--out", required=True)
    args = ap.parse_args()

    print(f"[*] loading: {args.src}")
    try:
        obj = torch.load(args.src, map_location="cpu", pickle_module=CustomPickleModule, weights_only=False)
    except TypeError:
        obj = torch.load(args.src, map_location="cpu", pickle_module=CustomPickleModule)

    mc = OmegaConf.create(walk(to_py(gf(obj, "model_config"))))
    tc = OmegaConf.create(walk(to_py(gf(obj, "tasks_config"))))
    sf(obj, "model_config", mc)
    sf(obj, "tasks_config", tc)

    cfg_py = OmegaConf.to_container(mc, resolve=False)
    bad = find_legacy(cfg_py)
    if bad:
        print("[FATAL] legacy paths still exist in model_config:")
        for p, v in bad[:20]:
            print(" ", p, "=", v)
        raise SystemExit(2)

    out = Path(args.out)
    out.parent.mkdir(parents=True, exist_ok=True)
    torch.save(obj, str(out))
    print("[OK] saved:", out)
    print("[OK] model _target_:", mc.get("_target_", "<none>"))
    print("[OK] backbone model:", mc.get("backbone", {}).get("model"))

if __name__ == "__main__":
    main()