CD-Models / utils /weight_downloader.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
9.84 kB
from __future__ import annotations
import argparse
import hashlib
import importlib
import json
import subprocess
import sys
import urllib.error
import urllib.request
from datetime import datetime, timezone
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
WEIGHT_REGISTRY = {
"vmamba_tiny": {
"filename": "vssmtiny_dp01_ckpt_epoch_292.pth",
"zenodo_record": "14037770",
"zenodo_url": "https://zenodo.org/records/14037770/files/vssmtiny_dp01_ckpt_epoch_292.pth",
"gdrive_id": "160PXughGMNZ1GyByspLFS68sfUdrQE2N",
"local_dir": "model_repos/ChangeMamba/pretrained_weight",
"sha256": None,
},
"vmamba_small": {
"filename": "vssmsmall_dp03_ckpt_epoch_238.pth",
"zenodo_record": "14037770",
"zenodo_url": "https://zenodo.org/records/14037770/files/vssmsmall_dp03_ckpt_epoch_238.pth",
"gdrive_id": "1dxHtFEgeJ9KL5WiLlvQOZK5jSEEd2Nmz",
"local_dir": "model_repos/ChangeMamba/pretrained_weight",
"sha256": None,
},
"vmamba_base": {
"filename": "vssmbase_dp06_ckpt_epoch_241.pth",
"zenodo_record": "14037770",
"zenodo_url": "https://zenodo.org/records/14037770/files/vssmbase_dp06_ckpt_epoch_241.pth",
"gdrive_id": "1kUHSBDoFvFG58EmwWurdSVZd8gyKWYfr",
"local_dir": "model_repos/ChangeMamba/pretrained_weight",
"sha256": None,
},
}
TIMM_MODEL_ALIASES = {
"efficientnet_b4": ("efficientnet_b4", "tf_efficientnet_b4", "tf_efficientnet_b4_ns"),
"mit_b0": ("mit_b0", "segformer_b0"),
"mit_b1": ("mit_b1", "segformer_b1"),
}
TORCHVISION_WEIGHT_ENUMS = {
"resnet18": "ResNet18_Weights",
"resnet50": "ResNet50_Weights",
"vgg16": "VGG16_Weights",
}
def _log_download(row: dict) -> None:
path = ROOT / "results" / "download_log.jsonl"
path.parent.mkdir(parents=True, exist_ok=True)
payload = {"timestamp_utc": datetime.now(timezone.utc).isoformat(), **row}
with path.open("a", encoding="utf-8") as f:
f.write(json.dumps(payload, sort_keys=True) + "\n")
def _sha256(filepath: str | Path) -> str:
h = hashlib.sha256()
with Path(filepath).open("rb") as f:
for block in iter(lambda: f.read(1024 * 1024), b""):
h.update(block)
return h.hexdigest()
def _manifest_path(local_dir: Path) -> Path:
return local_dir / "weights_manifest.json"
def _load_manifest(local_dir: Path) -> dict:
path = _manifest_path(local_dir)
if not path.exists():
return {}
with path.open("r", encoding="utf-8") as f:
return json.load(f)
def _save_manifest(local_dir: Path, manifest: dict) -> None:
with _manifest_path(local_dir).open("w", encoding="utf-8") as f:
json.dump(manifest, f, indent=2, sort_keys=True)
def _ensure_module(module: str, package: str | None = None):
try:
return importlib.import_module(module)
except ImportError:
subprocess.run([sys.executable, "-m", "pip", "install", package or module], check=True)
return importlib.import_module(module)
def ensure_weights(weight_key: str, verbose: bool = True) -> str:
if weight_key not in WEIGHT_REGISTRY:
raise KeyError(f"Unknown weight key {weight_key!r}. Known keys: {sorted(WEIGHT_REGISTRY)}")
spec = WEIGHT_REGISTRY[weight_key]
local_dir = ROOT / spec["local_dir"]
local_dir.mkdir(parents=True, exist_ok=True)
dest = local_dir / spec["filename"]
manifest = _load_manifest(local_dir)
if dest.exists():
expected = spec.get("sha256") or manifest.get(spec["filename"], {}).get("sha256")
if not _verify_sha256(str(dest), expected):
raise RuntimeError(f"SHA256 mismatch for existing weight: {dest}")
if verbose:
print(f"[WeightDownloader] Present: {dest}")
return str(dest.resolve())
ok = _download_from_zenodo(spec["zenodo_url"], str(dest))
if not ok:
ok = _download_from_gdrive(spec["gdrive_id"], str(dest))
if not ok or not dest.exists():
raise RuntimeError(f"Failed to download {weight_key} to {dest}")
digest = _sha256(dest)
manifest[spec["filename"]] = {"sha256": digest, "weight_key": weight_key}
_save_manifest(local_dir, manifest)
print(f"[WeightDownloader] Downloaded {dest}")
print(f"[WeightDownloader] SHA256 {digest}")
return str(dest.resolve())
def _download_from_zenodo(url: str, dest_path: str) -> bool:
try:
tqdm = _ensure_module("tqdm").tqdm
with urllib.request.urlopen(url) as response:
total = int(response.headers.get("Content-Length", "0"))
if getattr(response, "status", 200) != 200:
_log_download({"method": "zenodo", "url": url, "success": False, "status": response.status})
return False
with open(dest_path, "wb") as f, tqdm(total=total, unit="B", unit_scale=True, desc=Path(dest_path).name) as bar:
while True:
chunk = response.read(1024 * 1024)
if not chunk:
break
f.write(chunk)
bar.update(len(chunk))
_log_download({"method": "zenodo", "url": url, "success": True, "file_size": Path(dest_path).stat().st_size})
return True
except (urllib.error.URLError, OSError, subprocess.CalledProcessError) as exc:
_log_download({"method": "zenodo", "url": url, "success": False, "error": str(exc)})
return False
def _download_from_gdrive(file_id: str, dest_path: str) -> bool:
try:
gdown = _ensure_module("gdown")
url = f"https://drive.google.com/uc?id={file_id}"
result = gdown.download(url, dest_path, quiet=False)
ok = result is not None and Path(dest_path).exists()
_log_download({"method": "gdrive", "file_id": file_id, "success": ok, "file_size": Path(dest_path).stat().st_size if ok else 0})
return ok
except (OSError, subprocess.CalledProcessError) as exc:
_log_download({"method": "gdrive", "file_id": file_id, "success": False, "error": str(exc)})
return False
def _verify_sha256(filepath: str, expected: str | None) -> bool:
if expected is None:
return True
actual = _sha256(filepath)
if actual != expected:
print(f"[WeightDownloader] WARNING: SHA256 mismatch for {filepath}: expected {expected}, got {actual}")
return False
return True
def list_all_weights() -> None:
print("Weight Key | Filename | Local Path | Status")
print("--- | --- | --- | ---")
for key, spec in WEIGHT_REGISTRY.items():
path = ROOT / spec["local_dir"] / spec["filename"]
print(f"{key} | {spec['filename']} | {path} | {'present' if path.exists() else 'missing'}")
def ensure_timm_weight(model_name: str, required: bool = True) -> bool:
timm = _ensure_module("timm")
print(f"[WeightDownloader] Ensuring timm weights for: {model_name}")
candidates = TIMM_MODEL_ALIASES.get(model_name, (model_name,))
errors = []
for candidate in candidates:
try:
model = timm.create_model(candidate, pretrained=True, num_classes=0)
del model
_log_download({"method": "timm", "model_name": model_name, "resolved_model": candidate, "success": True})
print(f"[WeightDownloader] {model_name} weights ready via timm model {candidate}.")
return True
except RuntimeError as exc:
if "Unknown model" not in str(exc):
raise
errors.append(str(exc))
version = getattr(timm, "__version__", "unknown")
message = (
f"timm {version} does not provide {model_name} "
f"(tried: {', '.join(candidates)})."
)
_log_download({
"method": "timm",
"model_name": model_name,
"success": False,
"required": required,
"error": "; ".join(errors) or message,
})
if required:
raise RuntimeError(message)
print(f"[WeightDownloader] WARNING: {message} Skipping optional warmup.")
return False
def ensure_torchvision_weight(model_name: str) -> None:
tv_models = _ensure_module("torchvision.models", "torchvision")
print(f"[WeightDownloader] Ensuring torchvision weights for: {model_name}")
builder = getattr(tv_models, model_name)
enum_name = TORCHVISION_WEIGHT_ENUMS.get(
model_name,
"".join(part.capitalize() for part in model_name.split("_")) + "_Weights",
)
weights_enum = getattr(tv_models, enum_name, None)
if weights_enum is None:
builder(pretrained=True)
_log_download({"method": "torchvision", "model_name": model_name, "success": True, "weights": "pretrained=True"})
print(f"[WeightDownloader] {model_name} weights ready.")
return
weights = weights_enum.DEFAULT
builder(weights=weights)
_log_download({"method": "torchvision", "model_name": model_name, "success": True, "weights": str(weights)})
print(f"[WeightDownloader] {model_name} weights ready.")
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--prefetch-all", action="store_true")
parser.add_argument("--list", action="store_true")
args = parser.parse_args()
if args.list:
list_all_weights()
if args.prefetch_all:
for key in WEIGHT_REGISTRY:
ensure_weights(key, verbose=True)
ensure_timm_weight("efficientnet_b4")
ensure_timm_weight("mit_b0", required=False)
ensure_timm_weight("mit_b1", required=False)
ensure_torchvision_weight("resnet18")
ensure_torchvision_weight("resnet50")
ensure_torchvision_weight("vgg16")
print("[WeightDownloader] All weights prefetched successfully.")
return 0
if __name__ == "__main__":
raise SystemExit(main())