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()
|