项目文件夹

文件
wehub-resource-sync 593b94c120
pytest / Unit Tests (push) Has been cancelled
pytest / Integration (integration_tests_a) (push) Has been cancelled
pytest / Integration (integration_tests_b) (push) Has been cancelled
pytest / Integration (integration_tests_c) (push) Has been cancelled
pytest / Integration (integration_tests_d) (push) Has been cancelled
pytest / Integration (integration_tests_e) (push) Has been cancelled
pytest / Integration (integration_tests_f) (push) Has been cancelled
pytest / Integration (integration_tests_g) (push) Has been cancelled
pytest / Integration (integration_tests_h) (push) Has been cancelled
pytest / Integration (integration_tests_i) (push) Has been cancelled
pytest / Integration (integration_tests_j) (push) Has been cancelled
pytest / Distributed (distributed_a) (push) Has been cancelled
pytest / Distributed (distributed_b) (push) Has been cancelled
pytest / Distributed (distributed_c) (push) Has been cancelled
pytest / Distributed (distributed_d) (push) Has been cancelled
pytest / Distributed (distributed_e) (push) Has been cancelled
pytest / Distributed (distributed_f) (push) Has been cancelled
pytest / Minimal Install (push) Has been cancelled
pytest / Event File (push) Has been cancelled
pytest (slow) / py-slow (push) Has been cancelled
Publish JSON Schema / publish-schema (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 12:49:20 +08:00

60 行
1.7 KiB
Python

"""Tests for model export utilities."""
import os
import tempfile
import torch
import torch.nn as nn
from ludwig.utils.model_export import load_exported_model, ModelExporter
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(10, 5)
def forward(self, x):
return self.linear(x)
class TestModelExporter:
def test_export_safetensors(self):
model = SimpleModel()
exporter = ModelExporter(model)
with tempfile.TemporaryDirectory() as tmpdir:
path = exporter.export_safetensors(tmpdir)
assert os.path.exists(path)
assert path.endswith(".safetensors")
def test_export_torch(self):
model = SimpleModel()
exporter = ModelExporter(model)
sample = torch.randn(2, 10)
with tempfile.TemporaryDirectory() as tmpdir:
path = exporter.export_torch(tmpdir, sample)
assert os.path.exists(path)
def test_generate_sample_input_fallback(self):
model = SimpleModel()
exporter = ModelExporter(model)
sample = exporter._generate_sample_input()
assert "input" in sample
class TestLoadExportedModel:
def test_load_torchscript(self):
model = SimpleModel()
with tempfile.TemporaryDirectory() as tmpdir:
path = os.path.join(tmpdir, "model.pt")
traced = torch.jit.trace(model, torch.randn(2, 10))
traced.save(path)
loaded = load_exported_model(path)
assert loaded is not None
def test_unknown_format_raises(self):
import pytest
with pytest.raises(ValueError, match="Unknown model format"):
load_exported_model("model.xyz")