Exportarea modelelor cu TorchScript

Modele AI scalabile cu PyTorch Lightning

Sergiy Tkachuk

Director, GenAI Productivity

Ce este TorchScript?

  • Independent de Python

  • Eficient în producție

  • Exemple:

    • Implementare pe dispozitive mobile

Un telefon mobil

Modele AI scalabile cu PyTorch Lightning

Ce este TorchScript?

  • Independent de Python

  • Eficient în producție

  • Exemple:

    • Implementare pe dispozitive mobile
    • Inferență performantă în sisteme de producție

Un telefon mobil și un laptop cu angrenaje reprezentând un mediu de producție

Modele AI scalabile cu PyTorch Lightning

Conversia modelelor în TorchScript

Două metode de conversie:

  • torch.jit.trace: Urmărește execuția pe baza unor intrări exemplu
  • torch.jit.script: Compilează modelul analizând codul Python

Când se utilizează:

  1. Utilizați trace pentru modele simple
  2. Utilizați script pentru modele cu flux de control (ex.: bucle)
import torch
import torch.nn as nn

class SimpleModel(nn.Module): def forward(self, x): return x * 2
model = SimpleModel() scripted_model = torch.jit.script(model)
Modele AI scalabile cu PyTorch Lightning

Salvarea și încărcarea modelelor TorchScript

$$

  • Salvarea modelului:
    • torch.jit.save: Salvează modelul scriptat într-un fișier
  • Încărcarea modelului:
    • torch.jit.load: Încarcă modelul pentru inferență

$$

# Save the model
torch.jit.save(scripted_mod,"model.pt")

# Load the model
loaded_model=torch.jit.load("model.pt")
Modele AI scalabile cu PyTorch Lightning

Inferență cu TorchScript

  • Pași:
    • Încărcați modelul TorchScript
    • Transmiteți intrările modelului pentru predicții
    • Rezultatele sunt identice cu predicțiile PyTorch

Exemplu de intrare:

  • Tensor de intrare: [1.0, 2.0, 3.0]

Exemplu de ieșire:

  • Tensor de ieșire: [2.0, 4.0, 6.0]
# Perform inference
input_arr = [1.0, 2.0, 3.0]
input_tensor = torch.tensor(input_arr)

output = loaded_model(input_tensor) print(output)
Modele AI scalabile cu PyTorch Lightning

TorchScript pe scurt

$$

  • torch.jit.trace: Pentru modele statice
  • torch.jit.script: Gestionează fluxul de control dinamic
  • torch.jit.save: Salvează modelul scriptat
  • torch.jit.load: Reîncarcă pentru inferență
Modele AI scalabile cu PyTorch Lightning

Să exersăm!

Modele AI scalabile cu PyTorch Lightning

Preparing Video For Download...