0521_deep_hw / export_hf_model.py
yb1n's picture
Upload folder using huggingface_hub
beb73b4 verified
Raw
History Blame Contribute Delete
2.29 kB
import argparse
from pathlib import Path
from transformers import AutoTokenizer
from custom_text_classifier import CustomTextForSequenceClassification
from custom_text_config import CustomTextConfig
def parse_args():
parser = argparse.ArgumentParser(description="Export model files for Hugging Face")
parser.add_argument("--output-dir", default="./hf_model")
parser.add_argument("--backbone", default="distilbert-base-uncased")
return parser.parse_args()
def main():
args = parse_args()
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
config = CustomTextConfig(
backbone_name_or_path=args.backbone,
num_labels=2,
id2label={0: "negative", 1: "positive"},
label2id={"negative": 0, "positive": 1},
)
model = CustomTextForSequenceClassification(config)
tokenizer = AutoTokenizer.from_pretrained(args.backbone)
model.save_pretrained(str(output_dir))
tokenizer.save_pretrained(str(output_dir))
model_card = output_dir / "README.md"
model_card.write_text(
"\n".join(
[
"---",
"library_name: transformers",
"pipeline_tag: text-classification",
"tags:",
"- custom-code",
"- imdb",
"---",
"",
"# Custom Text Classifier",
"",
"This repository contains a custom `PreTrainedModel` sequence classifier.",
f"The backbone is `{args.backbone}` and the classification head has two labels.",
"",
"Load with:",
"",
"```python",
"from transformers import AutoModelForSequenceClassification, AutoTokenizer",
"",
"tokenizer = AutoTokenizer.from_pretrained(repo_id, trust_remote_code=True)",
"model = AutoModelForSequenceClassification.from_pretrained(",
" repo_id,",
" trust_remote_code=True,",
")",
"```",
"",
]
),
encoding="utf-8",
)
print("exported model folder:", output_dir)
if __name__ == "__main__":
main()