このプロジェクトについて

BurnはRustで書かれたテンソルライブラリ兼ディープラーニングフレームワークであり、数値計算、学習、推論を目的としている。掲げられている目標は、Pythonの学習スタックと別個のデプロイ用エンジンという一般的な分断を避けることである。同じモデルコードを学習してから本番環境で実行でき、プロジェクトではこれがオンデバイスのパーソナライゼーションや連合学習に有用だと説明されている。 設計と使い勝手 - PyTorchに似たAPIで動的シェイプとグラフを扱いながら、テンソル演算のストリームをJITコンパイルし、自動カーネル融合を実行する。 - インクリメンタルコンパイルが重視されており、READMEではリリースモードでもモデルコードの変更が5秒未満で再コンパイルされるとしている。 - モデルコードを変更せずにバックエンドを切り替え可能。1つのアプリケーションに複数のバックエンドを共存させられ、デバイスは実行時に`Device`で選択する。 バックエンド - GPU: CUDA、ROCm、Metal、Vulkan、WebGPU、LibTorch(0.22.0で非推奨)。対応状況はベンダーにより異なる。Nvidia(CUDA、Vulkan、WebGPU、LibTorch)、AMD(ROCm、Vulkan、WebGPU、LibTorch)、Apple(Metal、WebGPU、LibTorch)、IntelとQualcomm(Vulkan、WebGPU)、Wasm(WebGPU)。 - CPU: CubeCL CPUバックエンド、Flex、x86およびArm向けLibTorch。FlexはWasmとno_stdにも対応する。 - バックエンドデコレータ: Autodiffは任意のバックエンドに誤差逆伝播を追加する。Fusionは対応環境でカーネル融合を追加し、ファーストパーティのアクセラレーテッドバックエンドではデフォルトで有効。Remote(ベータ)はIrohまたはWebSocket経由のクライアント/サーバー実行をサポートし、分散計算に対応する。 学習と推論 - Ratatuiベースのターミナル学習ダッシュボードが学習・検証メトリクスをリアルタイム表示する。矢印キーで操作でき、クラッシュせずに学習ループから抜けられる。 - ONNXモデルはburn-onnxを通じてインポートし、Burn APIを使ってネイティブRustコードへ変換できるため、任意のBurnバックエンドで実行できる。READMEでは、このクレートは活発に開発中で演算子セットは限定的だと注記されている。 - PyTorchおよびSafetensors形式の重みをBurn定義のモデルに読み込める。 - 推論はWebAssembly経由でブラウザ内実行でき、Flex(CPU)またはWGPU(WebGPU)を使用する。MNISTと画像分類のデモがある。 - コアコンポーネントはベアメタル組み込み環境向けにno_stdをサポートする。現在no_stdで動作するのはFlexバックエンドのみ。 エコシステム - CubeCL: アクセラレーテッドバックエンドの背後にあるGPUコンピュート言語兼コンパイラで、単体でも利用可能。 - burn-onnx: ONNXインポート。burn-store: 重みの保存・読み込みとPyTorch/Safetensorsインポート。 - burn-vision、burn-rl、burn-datasetはそれぞれ画像処理、強化学習、データセット向け。 - modelsリポジトリには事前学習済みモデルとサンプルがある。burn-benchはバックエンドの経時的なベンチマーク用。 - コミュニティのクレート一覧には、データ読み込み(polars、arrow-rs、image、hf-hub)、トークナイズとNLP(tokenizers、rust-bert)、数値ライブラリ(ndarray、nalgebra)、古典的ML(linfa、smartcore)、推論ランタイム(candle、mistral.rs、ort、tract、wonnx)、LLM/RAGツール(rig、langchain-rust)、埋め込みとベクトル検索(fastembed、qdrant、lancedb)、コンピュータビジョン(kornia-rs)、シミュレーション(rapier)、可視化(rerun、plotters)が含まれる。 はじめ方 - Burn Bookが主要ドキュメントで、テンソル、モジュール、オプティマイザ、カスタムGPUカーネルを扱う。 - サンプルには、基本的なMNISTワークフロー、カスタム学習ループ、カスタムWGPUカーネル、CSVおよび画像データセット、回帰、カスタムレンダラー、ブラウザ推論デモ、PyTorch重みインポート、テキスト分類とテキスト生成、MNISTでのWGANが含まれる。 - ライセンスはMIT/Apache-2.0。