Download modeling.py from gradeeye/four-ch-morph: direct link, hf CLI and curl.
- Browser
- Download file 1.69 kB
-
https://huggingface.co/gradeeye/four-ch-morph/resolve/main/modeling.py
- Command line
-
hf download hf://gradeeye/four-ch-morph/modeling.py
-
curl -L -o modeling.py https://huggingface.co/gradeeye/four-ch-morph/resolve/main/modeling.py
1.69 kB
| """Minimal GradeEye classifier loader for Hugging Face Hub. | |
| Requires the GradeEye source package on PYTHONPATH plus torch, timm, and | |
| safetensors. The architecture is the same DRGradingModel used during training. | |
| EMA files are the primary weights and reproduce the paper evaluation protocol. | |
| """ | |
| from __future__ import annotations | |
| from pathlib import Path | |
| import sys | |
| import torch | |
| # For local source checkout usage. Users may instead install the GradeEye package. | |
| try: | |
| from src.models.dr_model import DRGradingModel | |
| except ImportError as exc: | |
| raise ImportError( | |
| "Install/clone GradeEye and make its repository root available on PYTHONPATH." | |
| ) from exc | |
| from safetensors.torch import load_file | |
| def load_model(weights_path: str | Path, config: dict, device: str = "cpu") -> DRGradingModel: | |
| """Instantiate DRGradingModel and strictly load a .safetensors state dict.""" | |
| model = DRGradingModel( | |
| pretrained=False, | |
| use_cbam=config["use_cbam"], | |
| cbam_num_stages=config["cbam_num_stages"], | |
| num_thresholds=config["num_thresholds"], | |
| head_hidden_dim=config["head_hidden_dim"], | |
| dropout=config["dropout"], | |
| output_mode="corn", | |
| arch=config["architecture"], | |
| in_chans=config["in_channels"], | |
| img_size=config["image_size"], | |
| ) | |
| state_dict = load_file(str(weights_path), device="cpu") | |
| result = model.load_state_dict(state_dict, strict=True) | |
| if result.missing_keys or result.unexpected_keys: | |
| raise RuntimeError( | |
| f"State-dict mismatch: missing={result.missing_keys}, " | |
| f"unexpected={result.unexpected_keys}" | |
| ) | |
| return model.to(device).eval() | |