FuXi-S2S / model /fuxi_s2s.py
yzt15806542928's picture
Upload folder using huggingface_hub
3fef26f verified
Raw
History Blame Contribute Delete
3.3 kB
"""Official FuXi-S2S ONNX adapter.
The public FuXi-S2S release contains an ONNX model and inference code. This
module exposes the official inference contract through ONNX Runtime.
"""
from pathlib import Path
import numpy as np
class FuXiS2SModel:
def __init__(self, model_path, device="cpu", providers=None):
self.model_path = Path(model_path)
self.device = device
self.providers = list(providers) if providers else self._default_providers(device)
try:
import onnxruntime as ort
except ImportError as exc:
raise RuntimeError("onnxruntime is required for official FuXi-S2S inference") from exc
if not self.model_path.is_file():
raise FileNotFoundError(f"FuXi-S2S ONNX model not found: {self.model_path}")
external_data_path = self.model_path.parent / "fuxi_s2s"
if not external_data_path.is_file():
raise FileNotFoundError(
f"FuXi-S2S ONNX external data not found: {external_data_path}. "
"The official archive requires both fuxi_s2s.onnx and fuxi_s2s."
)
available_providers = ort.get_available_providers()
missing_providers = [provider for provider in self.providers if provider not in available_providers]
if missing_providers:
raise RuntimeError(
f"ONNX Runtime provider(s) unavailable: {missing_providers}. "
f"Available providers: {available_providers}. "
"Install the matching ONNX Runtime build and update model.providers."
)
session_options = ort.SessionOptions()
if device == "dcu":
# CPU constant folding can warn on kernels that are intended for the
# accelerator. Let the configured DCU provider execute the graph.
session_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_DISABLE_ALL
self.session = ort.InferenceSession(
str(self.model_path),
sess_options=session_options,
providers=self.providers,
)
@staticmethod
def _default_providers(device):
if device == "cpu":
return ["CPUExecutionProvider"]
if device == "cuda":
return ["CUDAExecutionProvider", "CPUExecutionProvider"]
if device == "dcu":
raise ValueError(
"DCU inference requires an explicit ONNX Runtime provider in "
"model.providers, for example ['DCUExecutionProvider', "
"'CPUExecutionProvider']"
)
raise ValueError("device must be 'cpu', 'cuda', or 'dcu'")
@property
def input_names(self):
return [item.name for item in self.session.get_inputs()]
@property
def output_names(self):
return [item.name for item in self.session.get_outputs()]
def step(self, inputs):
payload = {}
for name in self.input_names:
if name not in inputs:
raise KeyError(f"Missing FuXi-S2S ONNX input: {name}")
payload[name] = np.asarray(inputs[name], dtype=np.float32)
outputs = self.session.run(None, payload)
return dict(zip(self.output_names, outputs))
def __call__(self, inputs):
return self.step(inputs)