РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
647 of 769 menu

Класс ScriptFunction

Класс ScriptFunction представляет собой скомпилированную версию функции Python, полученную с помощью декоратора @torch.jit.script. Объект этого класса содержит сериализованный граф вычислений, который может быть выполнен без доступа к исходному коду Python, что обеспечивает более высокую производительность и возможность сохранения на диск. Основное применение класса ScriptFunction - это создание моделей и функций в формате TorchScript для последующего использования в production-среде.

Синтаксис

import torch @torch.jit.script def function_name(params): # тело функции с операциями PyTorch return result

Пример

Давайте создадим простую скомпилированную функцию, которая вычисляет сумму квадратов элементов тензора:

import torch @torch.jit.script def sum_of_squares(x): return torch.sum(x * x) t = torch.tensor([1, 2, 3, 4, 5]) res = sum_of_squares(t) print(res)

Результат выполнения кода:

tensor(55)

Пример

Объект ScriptFunction можно сохранить на диск и загрузить в другом окружении:

import torch @torch.jit.script def add_tensors(a, b): return a + b t1 = torch.tensor([1, 2, 3]) t2 = torch.tensor([4, 5, 6]) res = add_tensors(t1, t2) print(res) torch.jit.save(add_tensors, 'add_tensors.pt')

Результат выполнения кода:

tensor([5, 7, 9])

Пример

Загруженный объект можно использовать как обычную функцию:

import torch loaded_func = torch.jit.load('add_tensors.pt') t3 = torch.tensor([10, 20, 30]) t4 = torch.tensor([1, 2, 3]) res = loaded_func(t3, t4) print(res)

Результат выполнения кода:

tensor([11, 22, 33])

Пример

Скомпилированная функция поддерживает использование условных операторов и циклов:

import torch @torch.jit.script def conditional_add(a, b, condition): if condition: return a + b else: return a - b t1 = torch.tensor([10, 20, 30]) t2 = torch.tensor([1, 2, 3]) cond = True res = conditional_add(t1, t2, cond) print(res)

Результат выполнения кода:

tensor([11, 22, 33])

Смотрите также

  • функцию jit.script,
    которая преобразует функцию Python в формат TorchScript
  • функцию jit.save,
    которая сохраняет объект TorchScript на диск
  • функцию jit.load,
    которая загружает объект TorchScript с диска
  • класс RecursiveScriptModule,
    который представляет скомпилированный модуль TorchScript
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить