Spouštění více modelů ML v řetězu

Windows ML podporuje vysoce výkonné zatížení a spouštění řetězů modelů pečlivou optimalizací cesty GPU. Řetězy modelů jsou definovány dvěma nebo více modely, které se spouštějí postupně, kde se výstupy jednoho modelu stanou vstupy do dalšího modelu dolů řetězu.

Abychom vysvětlili, jak efektivně zřetězovat modely pomocí Windows ML, použijeme jako příklad model FNS-Candy Style Transfer ONNX. Tento typ modelu najdete v ukázkové složce FNS-Candy Style Transfer na GitHubu.

Řekněme, že chceme spustit řetězec, který se skládá ze dvou instancí stejného modelu FNS-Candy, zde označovaného jako mosaic.onnx. Kód aplikace by předal image prvnímu modelu v řetězu, nechal ho vypočítat výstupy a pak předat tuto transformovanou image do jiné instance FNS-Candy, čímž se vytvoří konečná image.

Následující kroky ukazují, jak toho dosáhnout pomocí Windows ML.

Poznámka:

Ve skutečném scénáři byste pravděpodobně použili dva různé modely, ale to by mělo stačit k ilustraci konceptů.

  1. Nejprve načteme mosaic.onnx model, abychom ho mohli použít.
std::wstring filePath = L"path\\to\\mosaic.onnx"; 
LearningModel model = LearningModel::LoadFromFilePath(filePath);
string filePath = "path\\to\\mosaic.onnx";
LearningModel model = LearningModel.LoadFromFilePath(filePath);
  1. Pak vytvoříme dvě identické relace na výchozím GPU zařízení, kde použijeme stejný model jako vstupní parametr.
LearningModelSession session1(model, LearningModelDevice(LearningModelDeviceKind::DirectX));
LearningModelSession session2(model, LearningModelDevice(LearningModelDeviceKind::DirectX));
LearningModelSession session1 = 
  new LearningModelSession(model, new LearningModelDevice(LearningModelDeviceKind.DirectX));
LearningModelSession session2 = 
  new LearningModelSession(model, new LearningModelDevice(LearningModelDeviceKind.DirectX));

Poznámka:

Abyste mohli využívat výhody řetězení, musíte vytvořit identické GPU relace pro všechny vaše modely. To by vedlo k dalšímu přesunu dat z GPU do CPU, což by snížilo výkon.

  1. Následující řádky kódu vytvoří vazby pro každou relaci:
LearningModelBinding binding1(session1);
LearningModelBinding binding2(session2);
LearningModelBinding binding1 = new LearningModelBinding(session1);
LearningModelBinding binding2 = new LearningModelBinding(session2);
  1. V dalším kroku napojíme vstup pro náš první model. Předáme obrázek, který se nachází ve stejné cestě jako náš model. V tomto příkladu se image nazývá "fish_720.png".
//get the input descriptor
ILearningModelFeatureDescriptor input = model.InputFeatures().GetAt(0);
//load a SoftwareBitmap
hstring imagePath = L"path\\to\\fish_720.png";

// Get the image and bind it to the model's input
try
{
  StorageFile file = StorageFile::GetFileFromPathAsync(imagePath).get();
  IRandomAccessStream stream = file.OpenAsync(FileAccessMode::Read).get();
  BitmapDecoder decoder = BitmapDecoder::CreateAsync(stream).get();
  SoftwareBitmap softwareBitmap = decoder.GetSoftwareBitmapAsync().get();
  VideoFrame videoFrame = VideoFrame::CreateWithSoftwareBitmap(softwareBitmap);
  ImageFeatureValue image = ImageFeatureValue::CreateFromVideoFrame(videoFrame);
  binding1.Bind(input.Name(), image);
}
catch (...)
{
  printf("Failed to load/bind image\n");
}
//get the input descriptor
ILearningModelFeatureDescriptor input = model.InputFeatures[0];
//load a SoftwareBitmap
string imagePath = "path\\to\\fish_720.png";

// Get the image and bind it to the model's input
try
{
    StorageFile file = await StorageFile.GetFileFromPathAsync(imagePath);
    IRandomAccessStream stream = await file.OpenAsync(FileAccessMode.Read);
    BitmapDecoder decoder = await BitmapDecoder.CreateAsync(stream);
    SoftwareBitmap softwareBitmap = await decoder.GetSoftwareBitmapAsync();
    VideoFrame videoFrame = VideoFrame.CreateWithSoftwareBitmap(softwareBitmap);
    ImageFeatureValue image = ImageFeatureValue.CreateFromVideoFrame(videoFrame);
    binding1.Bind(input.Name, image);
}
catch
{
    Console.WriteLine("Failed to load/bind image");
}
  1. Aby další model v řetězci mohl používat výstupy z vyhodnocení prvního modelu, musíme vytvořit prázdný výstupní tenzor a svázat jeho výstup, abychom měli ukazatel pro propojení.
//get the output descriptor
ILearningModelFeatureDescriptor output = model.OutputFeatures().GetAt(0);
//create an empty output tensor 
std::vector<int64_t> shape = {1, 3, 720, 720};
TensorFloat outputValue = TensorFloat::Create(shape); 
//bind the (empty) output
binding1.Bind(output.Name(), outputValue);
//get the output descriptor
ILearningModelFeatureDescriptor output = model.OutputFeatures[0];
//create an empty output tensor 
List<long> shape = new List<long> { 1, 3, 720, 720 };
TensorFloat outputValue = TensorFloat.Create(shape);
//bind the (empty) output
binding1.Bind(output.Name, outputValue);

Poznámka:

Při vytváření vazby výstupu musíte použít datový typ TensorFloat . Po dokončení vyhodnocení prvního modelu se zabrání de-tensorizaci, což také eliminuje nutnost dalšího frontování GPU pro operace načítání a vazby u druhého modelu.

  1. Teď spustíme vyhodnocení prvního modelu a svážeme jeho výstupy se vstupem dalšího modelu:
//run session1 evaluation
session1.EvaluateAsync(binding1, L"");
//bind the output to the next model input
binding2.Bind(input.Name(), outputValue);
//run session2 evaluation
auto session2AsyncOp = session2.EvaluateAsync(binding2, L"");
//run session1 evaluation
await session1.EvaluateAsync(binding1, "");
//bind the output to the next model input
binding2.Bind(input.Name, outputValue);
//run session2 evaluation
LearningModelEvaluationResult results = await session2.EvaluateAsync(binding2, "");
  1. Nakonec načteme konečný výstup vytvořený po spuštění obou modelů pomocí následujícího řádku kódu.
auto finalOutput = session2AsyncOp.get().Outputs().First().Current().Value();
var finalOutput = results.Outputs.First().Value;

To je to! Vaše modely se teď mohou spouštět postupně a maximálně využívat dostupné prostředky GPU.

Poznámka:

Pomoc s Windows ML vám poskytnou následující zdroje:

  • Pokud chcete pokládat nebo odpovídat na technické otázky týkající se Windows ML, použijte značku windows-machine-learning ve službě Stack Overflow.
  • Pokud chcete nahlásit chybu, zapište prosím problém na našem GitHubu .