Spaces:
Sleeping
Sleeping
| import os | |
| import shutil | |
| from pathlib import Path | |
| from tempfile import TemporaryDirectory | |
| import gradio as gr | |
| def export_to_onnx( | |
| model_id: str, | |
| task: str, | |
| opset: int, | |
| push_to_hub: bool, | |
| repo_id: str, | |
| private_repo: bool, | |
| progress=gr.Progress(), | |
| ): | |
| model_id = model_id.strip() | |
| repo_id = repo_id.strip() | |
| if not model_id: | |
| raise gr.Error( | |
| "Enter a Hugging Face model ID, for example: " | |
| "distilbert/distilbert-base-uncased." | |
| ) | |
| if push_to_hub and not repo_id: | |
| raise gr.Error( | |
| "Enter a destination model repository, for example: " | |
| "your-username/my-model-onnx." | |
| ) | |
| if push_to_hub and not os.getenv("HF_TOKEN"): | |
| raise gr.Error( | |
| "This Space does not have an HF_TOKEN secret configured. " | |
| "Add a Hugging Face write token in Settings → Variables and secrets." | |
| ) | |
| if not task: | |
| task = "auto" | |
| try: | |
| from huggingface_hub import HfApi | |
| from optimum.exporters.onnx import main_export | |
| progress(0.05, desc="Preparing export...") | |
| with TemporaryDirectory() as temp_dir: | |
| temp_dir = Path(temp_dir) | |
| output_dir = temp_dir / "onnx" | |
| progress(0.15, desc="Downloading model and exporting to ONNX...") | |
| export_args = { | |
| "model_name_or_path": model_id, | |
| "output": output_dir, | |
| "task": task, | |
| "opset": int(opset), | |
| } | |
| if task == "auto": | |
| export_args.pop("task") | |
| main_export(**export_args) | |
| hub_url = None | |
| if push_to_hub: | |
| progress(0.80, desc="Creating or updating Hub repository...") | |
| api = HfApi(token=os.environ["HF_TOKEN"]) | |
| api.create_repo( | |
| repo_id=repo_id, | |
| repo_type="model", | |
| private=private_repo, | |
| exist_ok=True, | |
| ) | |
| api.upload_folder( | |
| folder_path=str(output_dir), | |
| repo_id=repo_id, | |
| repo_type="model", | |
| commit_message=( | |
| f"Export {model_id} to ONNX " | |
| f"(task={task}, opset={int(opset)})" | |
| ), | |
| ) | |
| hub_url = f"https://huggingface.co/{repo_id}" | |
| progress(0.90, desc="Creating download archive...") | |
| safe_name = model_id.replace("/", "--").replace("\\", "--") | |
| archive_base = temp_dir / f"{safe_name}-onnx" | |
| zip_path = shutil.make_archive( | |
| base_name=str(archive_base), | |
| format="zip", | |
| root_dir=str(output_dir), | |
| ) | |
| final_path = Path("/tmp") / f"{safe_name}-onnx.zip" | |
| shutil.copy2(zip_path, final_path) | |
| progress(1.0, desc="Finished.") | |
| message = f"Successfully exported `{model_id}` to ONNX." | |
| if hub_url: | |
| message += f"\n\nPushed to [{repo_id}]({hub_url})." | |
| return message, str(final_path) | |
| except Exception as error: | |
| raise gr.Error(f"Export failed: {type(error).__name__}: {error}") | |
| with gr.Blocks(title="Hugging Face to ONNX") as demo: | |
| gr.Markdown( | |
| """ | |
| # Hugging Face Model to ONNX | |
| Enter a public Hugging Face model ID, choose a task, then download the ONNX | |
| export or optionally push it to a Hugging Face model repository. | |
| **Examples:** `distilbert/distilbert-base-uncased`, | |
| `google-bert/bert-base-uncased`, or `gpt2` | |
| """ | |
| ) | |
| with gr.Row(): | |
| model_id = gr.Textbox( | |
| label="Hugging Face model ID", | |
| placeholder="distilbert/distilbert-base-uncased", | |
| scale=3, | |
| ) | |
| task = gr.Dropdown( | |
| label="Task", | |
| choices=[ | |
| "auto", | |
| "feature-extraction", | |
| "text-classification", | |
| "token-classification", | |
| "question-answering", | |
| "text-generation", | |
| "text2text-generation", | |
| "fill-mask", | |
| "image-classification", | |
| "audio-classification", | |
| ], | |
| value="auto", | |
| scale=2, | |
| ) | |
| opset = gr.Slider( | |
| label="ONNX opset", | |
| minimum=13, | |
| maximum=18, | |
| step=1, | |
| value=17, | |
| scale=1, | |
| ) | |
| gr.Markdown("## Optional Hub upload") | |
| with gr.Row(): | |
| push_to_hub = gr.Checkbox( | |
| label="Push ONNX export to Hugging Face Hub", | |
| value=False, | |
| scale=1, | |
| ) | |
| repo_id = gr.Textbox( | |
| label="Destination model repository", | |
| placeholder="your-username/model-name-onnx", | |
| scale=3, | |
| ) | |
| private_repo = gr.Checkbox( | |
| label="Private repository", | |
| value=False, | |
| scale=1, | |
| ) | |
| export_button = gr.Button("Export to ONNX", variant="primary") | |
| status = gr.Markdown() | |
| download = gr.File(label="Download ONNX ZIP") | |
| export_button.click( | |
| fn=export_to_onnx, | |
| inputs=[ | |
| model_id, | |
| task, | |
| opset, | |
| push_to_hub, | |
| repo_id, | |
| private_repo, | |
| ], | |
| outputs=[status, download], | |
| ) | |
| gr.Examples( | |
| examples=[ | |
| [ | |
| "distilbert/distilbert-base-uncased", | |
| "feature-extraction", | |
| 17, | |
| False, | |
| "", | |
| False, | |
| ], | |
| [ | |
| "google-bert/bert-base-uncased", | |
| "feature-extraction", | |
| 17, | |
| False, | |
| "", | |
| False, | |
| ], | |
| [ | |
| "distilbert/distilbert-base-cased-distilled-squad", | |
| "question-answering", | |
| 17, | |
| False, | |
| "", | |
| False, | |
| ], | |
| [ | |
| "gpt2", | |
| "text-generation", | |
| 17, | |
| False, | |
| "", | |
| False, | |
| ], | |
| ], | |
| inputs=[ | |
| model_id, | |
| task, | |
| opset, | |
| push_to_hub, | |
| repo_id, | |
| private_repo, | |
| ], | |
| ) | |
| if __name__ == "__main__": | |
| demo.launch() |