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