CELL-FM / upload_weights.py
BoHuangLab's picture
CELL-FM CondenSeq demo: sequence -> titration curve, AUC and AAC
e794567 verified
Raw History Blame Contribute Delete
3.6 kB
"""Push the three CELL-FM CondenSeq checkpoints to a HF model repo.
NOTE: never place credentials inside this directory -- it is uploaded verbatim to
the Space. Pass a token via the HF_TOKEN environment variable instead.
The Space pulls its weights from that repo at startup rather than carrying ~2.1 GB
of LFS itself. Run this once from the cluster:
python upload_weights.py --repo BoHuangLab/CELL-FM
Add --private to keep the weights unlisted, and --dry-run to see what would be
uploaded without touching the Hub.
"""
import argparse
import os
DEFAULT_SOURCES = {
# target name in the repo -> checkpoint on the cluster
"condenseq/cellfm_seq2img.bin": (
"/hpc/projects/group.huang/dihan.zheng/CELL-FM/"
"pretrain_condenseq/cellfm_seq2img/checkpoint-50000/pytorch_model.bin"
),
"condenseq/vae.bin": (
"/hpc/projects/group.huang/dihan.zheng/CELL-FM/"
"pretrain_condenseq/vae/checkpoint-50000/pytorch_model.bin"
),
"condenseq/vit_cls.bin": (
"/hpc/projects/group.huang/dihan.zheng/CELL-Diff2-Dev/"
"pretrain_condenseq_celldiff2_split/PT_CondenSeq_img_ViT_cls_R1/"
"checkpoint-10000/pytorch_model.bin"
),
}
CARD = """---
library_name: cell-fm
tags:
- biology
- microscopy
- protein
- condensate
- flow-matching
---
# CELL-FM weights
Checkpoints behind the [CELL-FM CondenSeq demo]({space_url}).
| File | Model | Source checkpoint |
|---|---|---|
| `condenseq/cellfm_seq2img.bin` | CELL-FM CS sequence-to-image generator (includes the ESM-C 600M encoder) | `pretrain_condenseq/cellfm_seq2img/checkpoint-50000` |
| `condenseq/vae.bin` | Image VAE, 160x160, 3 down blocks, 4 latent channels | `pretrain_condenseq/vae/checkpoint-50000` |
| `condenseq/vit_cls.bin` | ViT condensed/diffuse classifier, 2-channel 160x160 input | `PT_CondenSeq_img_ViT_cls_R1/checkpoint-10000` |
Hyperparameters are set in `pipeline.py` in the Space and mirror
`scripts/cell_fm_cs/evaluate_seq2img.sh` and
`scripts/vit_cls_condenseq_img/pretrain.sh` in the CELL-FM repository.
"""
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--repo", default="BoHuangLab/CELL-FM")
ap.add_argument("--space-url", default="https://huggingface.co/spaces/BoHuangLab/CELL-FM")
ap.add_argument("--private", action="store_true")
ap.add_argument("--dry-run", action="store_true")
args = ap.parse_args()
missing = [p for p in DEFAULT_SOURCES.values() if not os.path.exists(p)]
if missing:
raise SystemExit("missing checkpoint(s):\n " + "\n ".join(missing))
total = sum(os.path.getsize(p) for p in DEFAULT_SOURCES.values())
for name, path in DEFAULT_SOURCES.items():
print(f" {name:22s} {os.path.getsize(path)/1e9:5.2f} GB <- {path}")
print(f" {'total':22s} {total/1e9:5.2f} GB -> {args.repo}")
if args.dry_run:
print("\ndry run, nothing uploaded")
return
from huggingface_hub import HfApi
api = HfApi()
api.create_repo(args.repo, repo_type="model", private=args.private, exist_ok=True)
for name, path in DEFAULT_SOURCES.items():
print(f"uploading {name} ...")
api.upload_file(
path_or_fileobj=path,
path_in_repo=name,
repo_id=args.repo,
repo_type="model",
)
api.upload_file(
path_or_fileobj=CARD.format(space_url=args.space_url).encode(),
path_in_repo="README.md",
repo_id=args.repo,
repo_type="model",
)
print(f"\ndone: https://huggingface.co/{args.repo}")
if __name__ == "__main__":
main()