在C#环境中运行PyTorch模型,核心思路是利用TorchSharp这套.NET绑定库调用LibTorch原生能力。TorchSharp提供了与Python端近似的API,能够直接读取由PyTorch导出的序列化文件,并在内存中重建模型结构与参数。对于已经用Python训练完成的.pt或.pth文件,只要导出方式正确,C#侧并不需要重复编写网络定义。

一、PyTorch模型导出前的准备
在Python端保存模型时,常见做法有两种。一种是仅保存状态字典,即model.state_dict(),这种方式只保留权重,不保留网络结构;另一种是通过torch.jit.script或torch.jit.trace生成脚本化模型,将结构与参数一起固化到文件中。如果希望在C#中直接用TorchSharp读取并推理,推荐采用脚本化方式,因为C#侧无需重新声明网络层。
下面给出Python导出脚本化模型的示例。使用trace需要构造一个符合输入形状的示例张量,而script则直接分析代码路径。对于包含控制流的模型,script更稳妥。导出后得到的.pt文件可以在C#工程中直接加载。
import torch
import torch.nn as nn
class DemoNet(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 2)
def forward(self, x):
return self.fc(x)
net = DemoNet()
net.eval()
dummy = torch.randn(1, 10)
# 使用trace导出
traced = torch.jit.trace(net, dummy)
traced.save("demo_traced.pt")
# 或使用script导出
scripted = torch.jit.script(net)
scripted.save("demo_scripted.pt")
二、TorchSharp加载.pt文件的基础方式
TorchSharp中对应的加载方法是TorchSharp.Tensorflow?不对,应使用torch.jit.load。该方法位于TorchSharp命名空间下,返回的是一个可调用模块。加载时需注意文件路径使用绝对或相对可访问地址,且LibTorch运行时版本应与导出时PyTorch使用的LibTorch大版本兼容,否则会报序列化格式错误。
以下代码展示在C#控制台程序中读取脚本化模型,并在CPU上执行前向计算。注意输入张量的形状与Python端trace时的dummy保持一致,数据类型通常为float32。
using TorchSharp;
using static TorchSharp.torch;
class Program
{
static void Main()
{
// 加载脚本化模型
using var module = torch.jit.load("demo_scripted.pt");
// 构造输入张量,形状[1,10]
var input = torch.randn(1, 10);
// 前向推理
var output = module.forward(input);
Console.WriteLine(output);
}
}
如果模型导出时包含CUDA设备信息,而C#运行环境只有CPU,可在加载后调用module.to(device)进行设备迁移。TorchSharp通过torch.device方法声明目标设备,避免运行时报设备不匹配异常。
三、仅加载state_dict的权重恢复方案
当.pth文件只含state_dict时,C#侧必须先用C#重新定义相同结构的网络类,再读取权重字典逐项赋值。这种方式适合对模型结构有完全控制权的团队,但维护成本较高,一旦网络定义和Python端不一致就会加载失败。
下面示例展示如何定义线性网络并加载权重。TorchSharp的load方法可读取pth文件为字典对象,再通过module.load_state_dict映射。注意键名必须与Python端打印的state_dict键完全一致。
using TorchSharp;
using static TorchSharp.torch;
using TorchSharp.Modules;
class DemoNet : nn.Module<Tensor, Tensor>
{
private Linear fc;
public DemoNet() : base("DemoNet")
{
fc = Linear(10, 2);
RegisterComponents();
}
public override Tensor forward(Tensor x)
{
return fc.forward(x);
}
}
class Loader
{
static void Run()
{
var net = new DemoNet();
var state = torch.load("demo_weights.pth");
net.load_state_dict(state);
var inp = torch.randn(1, 10);
var outp = net.forward(inp);
Console.WriteLine(outp);
}
}
四、两种加载方式对比与避坑
脚本化加载优势在于结构内嵌、C#代码量少、跨语言一致性高;缺点是动态控制流复杂的模型可能trace不准。state_dict方式灵活、文件小,但要求双端网络定义同步,且容易因层名微调而失败。实际部署中,若模型来自第三方,优先协商提供scripted.pt。
常见错误包括:LibTorch版本不匹配导致无法读取、输入形状错误引发维度异常、未调用eval模式使Dropout等层干扰推理。建议在C#加载后立即用固定随机种子验证输出是否与Python端一致,从而确认加载链路正确。
| 加载方式 | 文件内容 | C#侧要求 | 适用场景 |
|---|---|---|---|
| jit.load | 结构与权重 | 无需定义网络 | 快速部署、跨语言 |
| load+state_dict | 仅权重 | 重定义网络 | 自研模型精细控制 |
五、设备与性能注意事项
在Windows或Linux服务器上运行TorchSharp,需安装对应LibTorch本地库。NuGet包TorchSharp默认带CPU版,若需CUDA应额外引用TorchSharp-cuda包。加载大模型时,建议显式指定device以减少默认设备探测开销。
推理批量处理时,可构造批维度输入并利用TorchSharp的no_grad上下文关闭梯度,降低内存占用。以下代码演示批量推理的基本写法。
using TorchSharp;
using static TorchSharp.torch;
void BatchInfer(string modelPath)
{
using var m = torch.jit.load(modelPath);
m.eval();
using (var noGrad = torch.no_grad())
{
var batch = torch.randn(32, 10);
var res = m.forward(batch);
Console.WriteLine(res.shape);
}
}
通过上述步骤,C#程序即可稳定读取PyTorch的.pt或.pth文件并完成推理任务。选择导出与加载方案时,应综合团队语言栈、模型复杂度和部署环境来决定。
TorchSharpPyTorch模型加载C#深度学习修改时间:2026-08-02 18:24:35