File size: 3,604 Bytes
e794567
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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()