lamossta commited on
Commit
9f3aa4a
·
1 Parent(s): 4820148

hf upload/download and onnx export

Browse files
src/models/hf_download.py ADDED
@@ -0,0 +1,68 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+ from pathlib import Path
4
+ from dotenv import load_dotenv
5
+ from huggingface_hub import hf_hub_download
6
+
7
+ load_dotenv()
8
+
9
+ MODELS_ROOT = Path("models")
10
+ VALID_MODES = ("marker", "qa_m", "qa_b", "fasttext")
11
+ REPO_PREFIX = "lamossta/distillbert"
12
+ FASTTEXT_REPO = "lamossta/fasttext_baseline"
13
+
14
+
15
+ def _repo_id(mode: str) -> str:
16
+ if mode == "fasttext":
17
+ return FASTTEXT_REPO
18
+ return f"{REPO_PREFIX}_{mode}"
19
+
20
+
21
+ def _model_filename(mode: str) -> str:
22
+ return "model.bin" if mode == "fasttext" else "model.onnx"
23
+
24
+
25
+ def download_model(mode: str, revision: str = "main") -> Path:
26
+ token = os.environ.get("HF_TOKEN")
27
+ repo_id = _repo_id(mode)
28
+ model_dir = MODELS_ROOT / mode
29
+ model_dir.mkdir(parents=True, exist_ok=True)
30
+ filename = _model_filename(mode)
31
+
32
+ hf_hub_download(
33
+ repo_id=repo_id,
34
+ filename=filename,
35
+ revision=revision,
36
+ token=token,
37
+ local_dir=str(model_dir),
38
+ )
39
+ print(f"Downloaded {repo_id}/{filename} -> {model_dir / filename}")
40
+
41
+ return model_dir
42
+
43
+
44
+ def download_all(revision: str = "main") -> dict[str, Path]:
45
+ downloaded = {}
46
+ for mode in VALID_MODES:
47
+ try:
48
+ downloaded[mode] = download_model(mode, revision)
49
+ except Exception as e:
50
+ print(f"Skipping '{mode}': {e}")
51
+ return downloaded
52
+
53
+
54
+ def main():
55
+ parser = argparse.ArgumentParser(description="Download ONNX models from Hugging Face")
56
+ parser.add_argument("--mode", default=None, choices=VALID_MODES, help="Single mode to download (default: all)")
57
+ parser.add_argument("--revision", default="main")
58
+ args = parser.parse_args()
59
+
60
+ if args.mode:
61
+ download_model(args.mode, args.revision)
62
+ else:
63
+ downloaded = download_all(args.revision)
64
+ print(f"Downloaded models: {downloaded}")
65
+
66
+
67
+ if __name__ == "__main__":
68
+ main()
src/models/hf_upload.py ADDED
@@ -0,0 +1,74 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+ from pathlib import Path
4
+ from dotenv import load_dotenv
5
+ from huggingface_hub import HfApi
6
+
7
+ load_dotenv()
8
+
9
+ MODELS_ROOT = Path("models")
10
+ VALID_MODES = ("marker", "qa_m", "qa_b", "fasttext")
11
+ REPO_PREFIX = "lamossta/distillbert"
12
+ FASTTEXT_REPO = "lamossta/fasttext_baseline"
13
+
14
+
15
+ def _repo_id(mode: str) -> str:
16
+ if mode == "fasttext":
17
+ return FASTTEXT_REPO
18
+ return f"{REPO_PREFIX}_{mode}"
19
+
20
+
21
+ def _model_filename(mode: str) -> str:
22
+ return "model.bin" if mode == "fasttext" else "model.onnx"
23
+
24
+
25
+ def upload_model(mode: str, revision: str = "main") -> None:
26
+ token = os.environ.get("HF_TOKEN")
27
+ api = HfApi(token=token)
28
+ model_dir = MODELS_ROOT / mode
29
+ repo_id = _repo_id(mode)
30
+ filename = _model_filename(mode)
31
+ model_path = model_dir / filename
32
+
33
+ if not model_path.exists():
34
+ raise FileNotFoundError(f"Model file not found: {model_path}")
35
+
36
+ api.create_repo(repo_id, exist_ok=True, token=token)
37
+
38
+ api.upload_file(
39
+ repo_id=repo_id,
40
+ path_or_fileobj=str(model_path),
41
+ path_in_repo=filename,
42
+ revision=revision,
43
+ token=token,
44
+ )
45
+ print(f"Uploaded '{mode}' model to {repo_id}/{filename}")
46
+
47
+
48
+ def upload_all(revision: str = "main") -> list[str]:
49
+ uploaded = []
50
+ for mode in VALID_MODES:
51
+ model_path = MODELS_ROOT / mode / _model_filename(mode)
52
+ if not model_path.exists():
53
+ print(f"Skipping '{mode}': {model_path} not found")
54
+ continue
55
+ upload_model(mode, revision)
56
+ uploaded.append(mode)
57
+ return uploaded
58
+
59
+
60
+ def main():
61
+ parser = argparse.ArgumentParser(description="Upload ONNX models to Hugging Face")
62
+ parser.add_argument("--mode", default=None, choices=VALID_MODES, help="Single mode to upload (default: all)")
63
+ parser.add_argument("--revision", default="main")
64
+ args = parser.parse_args()
65
+
66
+ if args.mode:
67
+ upload_model(args.mode, args.revision)
68
+ else:
69
+ uploaded = upload_all(args.revision)
70
+ print(f"Uploaded models: {uploaded}")
71
+
72
+
73
+ if __name__ == "__main__":
74
+ main()
src/models/serialize.py ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ from pathlib import Path
3
+ import torch
4
+ from transformers import AutoModelForSequenceClassification, AutoTokenizer
5
+
6
+
7
+ def export_to_onnx(
8
+ model_dir: str,
9
+ output_path: str | None = None,
10
+ max_len: int = 256,
11
+ opset_version: int = 17,
12
+ ) -> Path:
13
+ model_dir = Path(model_dir)
14
+ if output_path is None:
15
+ output_path = model_dir / "model.onnx"
16
+ else:
17
+ output_path = Path(output_path)
18
+
19
+ tokenizer = AutoTokenizer.from_pretrained(model_dir)
20
+ model = AutoModelForSequenceClassification.from_pretrained(model_dir)
21
+ model.eval()
22
+
23
+ dummy = tokenizer(
24
+ ["dummy input", "another dummy"],
25
+ max_length=max_len,
26
+ truncation=True,
27
+ padding="max_length",
28
+ return_tensors="pt",
29
+ )
30
+
31
+ batch = torch.export.Dim("batch", min=1, max=4096)
32
+ torch.onnx.export(
33
+ model,
34
+ (dummy["input_ids"], dummy["attention_mask"]),
35
+ str(output_path),
36
+ input_names=["input_ids", "attention_mask"],
37
+ output_names=["logits"],
38
+ dynamic_shapes={
39
+ "input_ids": {0: batch},
40
+ "attention_mask": {0: batch},
41
+ },
42
+ opset_version=opset_version,
43
+ dynamo=True,
44
+ external_data=False,
45
+ )
46
+
47
+ print(f"Exported ONNX model to '{output_path}' ({output_path.stat().st_size / 1024 / 1024:.1f} MB)")
48
+ return output_path
49
+
50
+
51
+ def main():
52
+ parser = argparse.ArgumentParser(description="Export model to ONNX")
53
+ parser.add_argument("--model-dir", required=True, help="Path to saved model directory")
54
+ parser.add_argument("--output", default=None, help="Output ONNX path (default: model_dir/model.onnx)")
55
+ parser.add_argument("--max-len", type=int, default=256)
56
+ parser.add_argument("--opset", type=int, default=17)
57
+ args = parser.parse_args()
58
+
59
+ export_to_onnx(args.model_dir, args.output, args.max_len, args.opset)
60
+
61
+
62
+ if __name__ == "__main__":
63
+ main()