ankane/torch.rb
Deep learning for Ruby, powered by LibTorch
Torch.rb – 为 Ruby 提供的深度学习
它是什么
- 一个 Ruby gem,封装了 Facebook 的 LibTorch(PyTorch 的 C++ 核心),让你可以在 Ruby 中编写和运行类似 PyTorch 的深度学习代码。
- 紧密模仿 PyTorch 的 Python API,同时进行了 Ruby 风格的调整(例如,使用
add!表示就地操作,tensor?用于布尔检查,以及使用Numo数组而非 NumPy)。
核心组件
- 张量操作 – 创建(
Torch.rand、Torch.zeros等)、算术运算、索引,以及与Numo::NArray之间的转换。 - 自动微分(Autograd) – 通过
requires_grad、backward和通过.grad访问梯度实现自动微分。 - 神经网络模块 – 通过继承
Torch::NN::Module定义模型,使用Conv2d、Linear等层,以及Torch::NN::F下的功能辅助函数。 - 优化器 – 例如
Torch::Optim::SGD,采用典型的zero_grad、step工作流。 - 保存/加载 – 使用
Torch.save/Torch.load保存和加载模型状态字典(与 PyTorch 文件兼容,但需注意少量 Python 端转换)。 - 设备支持 – 支持 CPU、CUDA GPU(
Torch::CUDA.available?、net.cuda)、Apple Silicon Metal(Torch::Backends::MPS)。
安装
- 下载适用于你平台的匹配 LibTorch 构建版本(仅 CPU 或支持 CUDA)。例如 macOS arm64 的示例:
curl -L https://download.pytorch.org/libtorch/cpu/libtorch-macos-arm64-2.14.0.zip > libtorch.zip unzip -q libtorch.zip - 告诉 gem LibTorch 的位置,并将其添加到你的
Gemfile中:bundle config set build.torch-rb --with-torch-dir=/path/to/libtorchgem "torch-rb" - 运行
bundle install。编译时间约为 5–10 分钟;不支持 Windows。
快速入门
- 参考
tutorials/blitz/README.md中的 60 分钟快速入门教程。 - 其他教程涵盖迁移学习、序列模型和词嵌入。
- 示例项目包括 MNIST 图像分类、MovieLens 协同过滤和 GAN。
典型工作流(Ruby 代码)
# 定义一个简单的 CNN
class MyNet < Torch::NN::Module
def initialize
super()
@conv1 = Torch::NN::Conv2d.new(1, 6, 3)
@conv2 = Torch::NN::Conv2d.new(6, 16, 3)
@fc1 = Torch::NN::Linear.new(16*6*6, 120)
@fc2 = Torch::NN::Linear.new(120, 84)
@fc3 = Torch::NN::Linear.new(84, 10)
end
def forward(x)
x = Torch::NN::F.max_pool2d(Torch::NN::F.relu(@conv1.call(x)), [2,2])
x = Torch::NN::F.max_pool2d(Torch::NN::F.relu(@conv2.call(x)), 2)
x = Torch.flatten(x, 1)
x = Torch::NN::F.relu(@fc1.call(x))
x = Torch::NN::F.relu(@fc2.call(x))
@fc3.call(x)
end
end
net = MyNet.new
input = Torch.randn(1,1,32,32)
output = net.call(input)
criterion = Torch::NN::MSELoss.new
target = Torch.randn(10).view(1,-1)
loss = criterion.call(output, target)
optimizer = Torch::Optim::SGD.new(net.parameters, lr: 0.01)
optimizer.zero_grad
loss.backward
optimizer.step
设备处理
if Torch::CUDA.available?
net.cuda # 将模型移动到 GPU
input = input.cuda
end
或在 Apple Silicon 上:
if Torch::Backends::MPS.available?
device = Torch.device('mps')
net.to(device)
end
生态系统
- 针对特定领域的配套 gem:
torchvision-ruby、torchtext-ruby、torchaudio-ruby、torchcodec-ruby、torchrec-ruby、torchdata-ruby。 - 高层级工具:
transformers-ruby(Transformer 模型)和safetensors-ruby(高效张量存储)。
为何重要
- 让 Ruby 开发者无需切换语言即可尝试现代深度学习模型。
- 利用与 PyTorch 相同的高性能 C++ 后端,使模型在 CPU 或 GPU 上以原生速度运行。
- 通过熟悉的 API 保持 Ruby 社区与快速演进的 PyTorch 生态系统同步。
资源
- 构建状态徽章显示在多个平台上的 CI 测试情况。
- 提供详细的变更日志、贡献指南以及 GPU 测试脚本(AWS Deep Learning AMI)。
- Torch.rb 将 PyTorch 的强大功能带到了 Ruby,提供了一个功能完整、支持 GPU 加速的深度学习栈,让 Ruby 程序员感觉如同原生使用一般。*
相关
- 项目
- 项目
- 项目
- 项目