File size: 5,300 Bytes
bd659a9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import argparse
import json
import shutil
import sys
from pathlib import Path

import torch
from huggingface_hub import snapshot_download

from local_config import load_local_api_config, repo_root, resolve_repo_path


ROOT = repo_root()
if str(ROOT) not in sys.path:
    sys.path.insert(0, str(ROOT))

from src.models.local_store import expected_local_model_path


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Download a Hugging Face text model once into Fuse-MD's local_models store."
    )
    parser.add_argument(
        "--model",
        help="Hugging Face model id to download, for example VishnuPJ/MalayaLLM_7B_Base.",
    )
    parser.add_argument(
        "--checkpoint",
        help="Checkpoint path to inspect and infer the required text model.",
    )
    parser.add_argument(
        "--root",
        help="Override the local model root. Defaults to api/local_config.py.",
    )
    parser.add_argument(
        "--force",
        action="store_true",
        help="Re-download into the target folder even if it already exists.",
    )
    return parser.parse_args()


def infer_text_model_from_checkpoint(checkpoint_path: Path) -> str:
    checkpoint = torch.load(checkpoint_path, map_location="cpu")
    checkpoint_cfg = checkpoint.get("config", {})
    if isinstance(checkpoint_cfg.get("model"), dict):
        checkpoint_cfg = checkpoint_cfg["model"]

    text_model = checkpoint_cfg.get("text_model")
    if not text_model:
        raise KeyError(f"Checkpoint does not define a text model: {checkpoint_path}")
    return str(text_model)


def choose_model_id(args: argparse.Namespace) -> str:
    if args.model:
        return str(args.model)

    if args.checkpoint:
        checkpoint_path = resolve_repo_path(args.checkpoint)
        if not checkpoint_path.exists():
            raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
        return infer_text_model_from_checkpoint(checkpoint_path)

    config = load_local_api_config()
    if not config.checkpoint_path.exists():
        raise FileNotFoundError(
            "No model id was provided and the default checkpoint is missing. "
            "Pass --model or --checkpoint."
        )
    return infer_text_model_from_checkpoint(config.checkpoint_path)


def resolved_root_arg(args: argparse.Namespace) -> str:
    if args.root:
        return str(resolve_repo_path(args.root))
    return str(load_local_api_config().local_model_root)


def verify_model_files(target_path: Path) -> None:
    required_files = [target_path / "config.json", target_path / "tokenizer_config.json"]
    missing_files = [path.name for path in required_files if not path.exists()]
    if missing_files:
        raise FileNotFoundError(
            f"Downloaded model is incomplete at {target_path}. Missing: {', '.join(missing_files)}"
        )

    has_weights = any(
        (target_path / file_name).exists()
        for file_name in (
            "model.safetensors",
            "model.safetensors.index.json",
            "pytorch_model.bin",
            "pytorch_model.bin.index.json",
        )
    )
    if not has_weights:
        raise FileNotFoundError(
            f"Downloaded model is incomplete at {target_path}. Missing model weights."
        )

    for index_name in ("pytorch_model.bin.index.json", "model.safetensors.index.json"):
        index_path = target_path / index_name
        if not index_path.exists():
            continue

        with open(index_path, "r", encoding="utf-8") as file:
            index_payload = json.load(file)

        weight_map = index_payload.get("weight_map", {})
        shard_names = sorted(set(str(name) for name in weight_map.values()))
        missing_shards = [name for name in shard_names if not (target_path / name).exists()]
        if missing_shards:
            raise FileNotFoundError(
                f"Downloaded model is incomplete at {target_path}. Missing shard files: "
                f"{', '.join(missing_shards)}"
            )


def safe_remove_tree(target_path: Path, root_path: Path) -> None:
    resolved_target = target_path.resolve()
    resolved_root = root_path.resolve()
    if resolved_target == resolved_root or resolved_root not in resolved_target.parents:
        raise ValueError(f"Refusing to remove path outside local model root: {resolved_target}")
    shutil.rmtree(resolved_target)


def main() -> int:
    args = parse_args()
    model_id = choose_model_id(args)
    root_arg = resolved_root_arg(args)
    root_path = Path(root_arg)
    target_path = expected_local_model_path(model_id, root_arg)
    root_path.mkdir(parents=True, exist_ok=True)

    if target_path.exists() and args.force:
        safe_remove_tree(target_path, root_path)

    if target_path.exists():
        verify_model_files(target_path)
        print(f"Local model already exists: {target_path}")
        print("Use --force to re-download it.")
        return 0

    print(f"Downloading model: {model_id}")
    print(f"Target folder: {target_path}")
    snapshot_download(repo_id=model_id, local_dir=str(target_path))
    verify_model_files(target_path)
    print("Local model download complete.")
    print(f"Saved to: {target_path}")
    return 0


if __name__ == "__main__":
    sys.exit(main())