File size: 4,915 Bytes
59630ba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
"""
This repo is forked from [Boyuan Chen](https://boyuan.space/)'s research 
template [repo](https://github.com/buoyancy99/research-template). 
By its MIT license, you must keep the above sentence in `README.md` 
and the `LICENSE` file to credit the author.
"""

from typing import Literal, Optional, Tuple
import string
import random
from pathlib import Path
from omegaconf import DictConfig
import wandb
from utils.print_utils import cyan
from .huggingface_utils import download_from_hf


def is_run_id(run_id: str) -> bool:
    """Check if a string is a run ID."""
    return len(run_id) == 8 and run_id.isalnum()


def generate_run_id() -> str:
    """Generate a random 8-character alphanumeric string."""
    chars = string.ascii_lowercase + string.digits
    return "".join(random.choice(chars) for _ in range(8))


def generate_unexisting_run_id(entity: str, project: str) -> str:
    """Generate a random 8-character alphanumeric string that does not exist in the project."""
    api = wandb.Api()
    runs = api.runs(f"{entity}/{project}")
    existing_ids = {run.id for run in runs}
    while True:
        run_id = generate_run_id()
        if run_id not in existing_ids:
            return run_id


def parse_load(load: str) -> Tuple[Optional[str], Optional[str]]:
    """
    Parse load into run_id and download option.
    (for load=xxxxxxxx in configurations)
    - If load_id is a run_id, return the run_id and None.
    - If load_id is of the form run_id:option, return run_id and option.
    - Otherwise, return None, None.
    """
    split = load.split(":")
    if 1 <= len(split) <= 2 and is_run_id(split[0]):
        return split[0], split[1] if len(split) == 2 else None
    return None, None


def version_to_int(artifact) -> int:
    """Convert versions of the form vX to X. For example, v12 to 12."""
    return int(artifact.version[1:])


def is_existing_run(run_path: str) -> bool:
    """Check if a run exists."""
    api = wandb.Api()
    try:
        _ = api.run(run_path)
        return True
    except wandb.errors.CommError:
        return False
    return False


def has_checkpoint(run_path: str) -> bool:
    """Check if a run has a committed model checkpoint."""
    api = wandb.Api()
    try:
        run = api.run(run_path)
        for artifact in run.logged_artifacts():
            if artifact.type == "model" and artifact.state == "COMMITTED":
                return True
        return False
    except wandb.errors.CommError:
        return False
    return False


def download_checkpoint(
    run_path: str, download_dir: Path, option: Literal["latest", "best"] = "latest"
) -> Path:
    api = wandb.Api()
    run = api.run(run_path)

    # Find the latest saved model checkpoint.
    checkpoint = None
    for artifact in run.logged_artifacts():
        if artifact.type != "model" or artifact.state != "COMMITTED":
            continue

        if option in artifact.aliases or option == artifact.version:
            checkpoint = artifact
            break

    if checkpoint is None:
        print(f"No {option} model checkpoint found in {run_path}.")

    # Download the checkpoint.
    download_dir.mkdir(exist_ok=True, parents=True)
    root = download_dir / run_path
    checkpoint.download(root=root)
    return root / "model.ckpt"


def download_pretrained(
    name: str,
) -> str:
    """
    Download a pretrained model from the DFoT Hugging Face model hub.
    Set is_full to True to download the full model
    (including optimizer states and non-EMA weights).
    """
    prefix, name = name.split(":")
    download_from_hf(filename="config.json")
    return download_from_hf(filename=f"{prefix}_models/{name}")


def is_wandb_run_path(run_path: str) -> bool:
    split = run_path.split("/")
    return len(split) == 3 and is_run_id(split[-1])


def is_hf_path(path: str) -> bool:
    return path.startswith("pretrained:") or path.startswith("full:")


def download_vae_checkpoints(
    cfg: DictConfig,
):
    pretrained_paths = []
    vae = cfg.algorithm.get("vae", None)
    if vae and vae.get("pretrained_path", None):
        pretrained_paths.append(vae.pretrained_path)

    pretrained_path = cfg.algorithm.get("pretrained_path", None)
    if pretrained_path:
        pretrained_paths.append(pretrained_path)

    wandb_pretrained_paths = [
        path for path in pretrained_paths if is_wandb_run_path(path)
    ]
    hf_pretrained_paths = [path for path in pretrained_paths if is_hf_path(path)]

    for path in wandb_pretrained_paths:
        print(cyan("Downloading pretrained VAE from Wandb:"), path)
        download_checkpoint(path, Path("outputs/downloaded"), option="best")

    for path in hf_pretrained_paths:
        print(cyan("Downloading pretrained VAE from Hugging Face:"), path)
        download_pretrained(path)


def wandb_to_local_path(run_path: str) -> Path:
    return Path("outputs/downloaded") / run_path / "model.ckpt"