Функция lerp
Функция lerp выполняет линейную интерполяцию между двумя тензорами.
Она вычисляет результат по формуле: result = start + weight * (end - start).
Первым параметром передаётся начальный тензор start,
вторым - конечный тензор end,
третьим - вес weight, который может быть числом или тензором.
Все тензоры должны иметь одинаковую форму или быть вещаемыми.
Синтаксис
torch.lerp(start, end, weight)
Пример
Давайте выполним линейную интерполяцию между двумя тензорами с весом 0.5:
import torch
start = torch.tensor([1.0, 2.0, 3.0])
end = torch.tensor([5.0, 6.0, 7.0])
res = torch.lerp(start, end, 0.5)
print(res)
Результат выполнения кода:
tensor([3., 4., 5.])
Пример
Теперь используем вес в виде тензора для поэлементной интерполяции:
import torch
start = torch.tensor([1.0, 2.0, 3.0])
end = torch.tensor([5.0, 6.0, 7.0])
weight = torch.tensor([0.1, 0.5, 0.9])
res = torch.lerp(start, end, weight)
print(res)
Результат выполнения кода:
tensor([1.4, 4.0, 6.6])
Пример
Используем функцию lerp для двумерных тензоров:
import torch
start = torch.tensor([
[1.0, 2.0],
[3.0, 4.0],
])
end = torch.tensor([
[5.0, 6.0],
[7.0, 8.0],
])
res = torch.lerp(start, end, 0.3)
print(res)
Результат выполнения кода:
tensor([
[2.2, 3.2],
[4.2, 5.2],
])
Пример
Вес может быть скаляром, а тензоры могут иметь разные размерности благодаря механизму вещания:
import torch
start = torch.tensor([1.0, 2.0, 3.0])
end = torch.tensor([5.0, 6.0, 7.0])
res = torch.lerp(start, end, 0.75)
print(res)
Результат выполнения кода:
tensor([4., 5., 6.])