TensorVizion's picture
Update app.py
00598ab verified
Raw History Blame Contribute Delete
12.3 kB
import os
import shutil
import subprocess
import tempfile
from pathlib import Path
import gradio as gr
import numpy as np
import torch
from safetensors.torch import load_file, save_file
try:
import h5py
H5PY_AVAILABLE = True
except ImportError:
H5PY_AVAILABLE = False
PYTORCH_EXTENSIONS = {".pt", ".pth", ".bin", ".ckpt"}
HDF5_EXTENSIONS = {".h5", ".hdf5"}
SUPPORTED_INPUTS = PYTORCH_EXTENSIONS | {".safetensors", ".npz"} | HDF5_EXTENSIONS
ONNX_TASKS = [
"auto",
"text-classification",
"token-classification",
"question-answering",
"feature-extraction",
"text-generation",
"text2text-generation",
"image-classification",
"image-to-text",
"semantic-segmentation",
]
def get_file_path(uploaded_file):
"""Works with Gradio versions that return either a filepath string or file object."""
if isinstance(uploaded_file, str):
return uploaded_file
return uploaded_file.name
def normalize_state_dict(data):
"""Extract a flat tensor dictionary from common checkpoint structures."""
if not isinstance(data, dict):
raise ValueError(
"The file does not contain a state dictionary. "
"Full serialized model objects are not supported."
)
for key in ("state_dict", "model_state_dict", "model"):
if key in data and isinstance(data[key], dict):
data = data[key]
break
result = {}
skipped = []
for key, value in data.items():
if isinstance(value, torch.Tensor):
result[str(key)] = value.detach().cpu().contiguous()
elif isinstance(value, np.ndarray):
result[str(key)] = torch.from_numpy(value).detach().cpu().contiguous()
else:
try:
result[str(key)] = torch.as_tensor(value).detach().cpu().contiguous()
except Exception:
skipped.append(str(key))
if not result:
raise ValueError("No tensor weights were found in this file.")
return result, skipped
def read_hdf5_datasets(group, prefix=""):
"""Read nested HDF5 datasets while preserving their group paths."""
state_dict = {}
for name, item in group.items():
key = f"{prefix}/{name}" if prefix else name
if isinstance(item, h5py.Dataset):
state_dict[key] = torch.from_numpy(np.asarray(item))
elif isinstance(item, h5py.Group):
state_dict.update(read_hdf5_datasets(item, key))
return state_dict
def convert_model(uploaded_file, target_format):
if uploaded_file is None:
return None, "❌ Please upload a model-weight file."
input_path = get_file_path(uploaded_file)
input_ext = Path(input_path).suffix.lower()
if input_ext not in SUPPORTED_INPUTS:
return None, f"❌ Unsupported input format: `{input_ext or 'no extension'}`."
if target_format == ".onnx":
return None, (
"⚠️ ONNX cannot be created from a weights file alone. "
"Use the **Hugging Face → ONNX** tab and enter a Hugging Face model ID. "
"ONNX needs the model architecture and example inputs."
)
if target_format in HDF5_EXTENSIONS and not H5PY_AVAILABLE:
return None, "❌ HDF5 support is unavailable. Add `h5py` to `requirements.txt`."
try:
if input_ext in PYTORCH_EXTENSIONS:
try:
data = torch.load(input_path, map_location="cpu", weights_only=True)
except Exception:
return None, (
"❌ This PyTorch checkpoint cannot be safely loaded with "
"`weights_only=True`. Do not upload untrusted pickle-based "
"files; they can execute code when deserialized."
)
state_dict, skipped = normalize_state_dict(data)
elif input_ext == ".safetensors":
state_dict, skipped = normalize_state_dict(load_file(input_path))
elif input_ext == ".npz":
with np.load(input_path, allow_pickle=False) as archive:
data = {key: archive[key] for key in archive.files}
state_dict, skipped = normalize_state_dict(data)
elif input_ext in HDF5_EXTENSIONS:
if not H5PY_AVAILABLE:
return None, "❌ HDF5 support is unavailable. Add `h5py` to `requirements.txt`."
with h5py.File(input_path, "r") as h5_file:
data = read_hdf5_datasets(h5_file)
state_dict, skipped = normalize_state_dict(data)
else:
return None, f"❌ Unsupported input format: `{input_ext}`."
except Exception as error:
return None, f"❌ Could not load file: `{type(error).__name__}: {error}`"
try:
# Do NOT use TemporaryDirectory here: Gradio needs the returned file to persist.
fd, output_path = tempfile.mkstemp(
prefix="converted_model_",
suffix=target_format,
)
os.close(fd)
if target_format in PYTORCH_EXTENSIONS:
torch.save(state_dict, output_path)
elif target_format == ".safetensors":
save_file(state_dict, output_path)
elif target_format == ".npz":
np.savez(
output_path,
**{key: value.cpu().numpy() for key, value in state_dict.items()},
)
elif target_format in HDF5_EXTENSIONS:
with h5py.File(output_path, "w") as h5_file:
for key, value in state_dict.items():
h5_file.create_dataset(key, data=value.cpu().numpy())
else:
os.unlink(output_path)
return None, f"❌ Unsupported target format: `{target_format}`."
message = f"✅ Converted {len(state_dict)} tensor(s) to `{target_format}`."
if skipped:
message += f" Skipped {len(skipped)} non-tensor value(s)."
return output_path, message
except Exception as error:
if "output_path" in locals() and os.path.exists(output_path):
os.unlink(output_path)
return None, f"❌ Could not save file: `{type(error).__name__}: {error}`"
def export_hf_model_to_onnx(model_id, task, optimization):
"""
Export a Hugging Face Hub model to ONNX with Optimum.
Example model IDs:
- distilbert/distilbert-base-uncased-finetuned-sst-2-english
- google-bert/bert-base-uncased
"""
if not model_id or not model_id.strip():
return None, "❌ Enter a Hugging Face model ID."
model_id = model_id.strip()
output_dir = tempfile.mkdtemp(prefix="optimum_onnx_")
command = [
"optimum-cli",
"export",
"onnx",
"--model",
model_id,
]
if task and task != "auto":
command.extend(["--task", task])
if optimization and optimization != "None":
command.extend(["--optimize", optimization])
command.append(output_dir)
try:
result = subprocess.run(
command,
capture_output=True,
text=True,
timeout=1800,
check=False,
)
if result.returncode != 0:
error_message = result.stderr.strip() or result.stdout.strip()
shutil.rmtree(output_dir, ignore_errors=True)
return None, (
"❌ ONNX export failed.\n\n"
f"`{error_message[-3500:]}`"
)
safe_name = model_id.replace("/", "_").replace("\\", "_")
archive_base = os.path.join(
tempfile.gettempdir(),
f"{safe_name}_onnx_export",
)
zip_path = shutil.make_archive(archive_base, "zip", output_dir)
onnx_files = list(Path(output_dir).rglob("*.onnx"))
onnx_count = len(onnx_files)
shutil.rmtree(output_dir, ignore_errors=True)
return zip_path, (
f"✅ ONNX export finished. Found {onnx_count} ONNX file(s). "
"Download the ZIP and extract it before using the model."
)
except subprocess.TimeoutExpired:
shutil.rmtree(output_dir, ignore_errors=True)
return None, (
"❌ Export timed out after 30 minutes. "
"Try a smaller model or use more powerful Space hardware."
)
except FileNotFoundError:
shutil.rmtree(output_dir, ignore_errors=True)
return None, (
"❌ `optimum-cli` was not found. Add "
"`optimum-onnx[onnxruntime]` to `requirements.txt`."
)
except Exception as error:
shutil.rmtree(output_dir, ignore_errors=True)
return None, f"❌ ONNX export error: `{type(error).__name__}: {error}`"
with gr.Blocks(theme=gr.themes.Soft(), title="Conversion Station") as demo:
gr.Markdown(
"""
# 🔄 Universal Model Weight Converter
Convert compatible **weight dictionaries** among PyTorch, SafeTensors, NumPy NPZ, and HDF5.
- **Weight inputs:** `.pt`, `.pth`, `.bin`, `.ckpt`, `.safetensors`, `.npz`, `.h5`, `.hdf5`
- **Weight outputs:** `.safetensors`, `.pt`, `.pth`, `.bin`, `.ckpt`, `.npz`, `.h5`, `.hdf5`
- **ONNX:** Use the separate Hugging Face → ONNX tab. ONNX requires a full model architecture, not only a weights file.
- **Security:** Only upload checkpoint files you trust. Pickle-based PyTorch files can be unsafe.
"""
)
with gr.Tab("Weight Converter"):
with gr.Row():
with gr.Column():
input_file = gr.File(
label="Upload model weights",
file_types=sorted(SUPPORTED_INPUTS),
type="filepath",
)
target_format = gr.Dropdown(
choices=[
".safetensors",
".pt",
".pth",
".bin",
".ckpt",
".npz",
".h5",
".hdf5",
".onnx",
],
value=".safetensors",
label="Target format",
)
convert_button = gr.Button("🚀 Convert model", variant="primary")
with gr.Column():
status_text = gr.Textbox(
label="Status",
interactive=False,
lines=4,
)
output_file = gr.File(label="Download converted model")
convert_button.click(
fn=convert_model,
inputs=[input_file, target_format],
outputs=[output_file, status_text],
)
with gr.Tab("Hugging Face → ONNX"):
gr.Markdown(
"""
Enter a public Hugging Face Transformers model ID to export it to ONNX.
Examples:
- `distilbert/distilbert-base-uncased-finetuned-sst-2-english`
- `google-bert/bert-base-uncased`
- `sentence-transformers/all-MiniLM-L6-v2`
The download is a ZIP because ONNX exports can include several ONNX graphs and supporting configuration files.
"""
)
with gr.Row():
onnx_model_id = gr.Textbox(
label="Hugging Face model ID",
placeholder="distilbert/distilbert-base-uncased-finetuned-sst-2-english",
)
onnx_task = gr.Dropdown(
choices=ONNX_TASKS,
value="auto",
label="Model task",
)
onnx_optimization = gr.Dropdown(
choices=["None", "O1", "O2", "O3", "O4"],
value="O2",
label="Optimization level",
)
onnx_export_button = gr.Button(
"⚡ Export Hugging Face Model to ONNX",
variant="primary",
)
onnx_status = gr.Textbox(
label="ONNX export status",
interactive=False,
lines=5,
)
onnx_output_file = gr.File(
label="Download ONNX export ZIP",
file_types=[".zip"],
)
onnx_export_button.click(
fn=export_hf_model_to_onnx,
inputs=[onnx_model_id, onnx_task, onnx_optimization],
outputs=[onnx_output_file, onnx_status],
)
if __name__ == "__main__":
demo.launch()