From d95f8ea6c0b388ecbf0b0c6abf68f336535cfcf9 Mon Sep 17 00:00:00 2001 From: Martin Bukowski Date: Tue, 2 Jan 2024 16:14:13 -0600 Subject: [PATCH] more cuda --- merge/mergeutil.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/merge/mergeutil.py b/merge/mergeutil.py index 9cc7357..dd878cc 100644 --- a/merge/mergeutil.py +++ b/merge/mergeutil.py @@ -146,6 +146,7 @@ def merge_tensors_cyclic(v0: torch.Tensor, v1: torch.Tensor, t: float) -> torch. # Model 1 > Model 2 > Model 1, with t defining the peak of the gradient along the tensor's width def merge_tensors_gradient(v0: torch.Tensor, v1: torch.Tensor, t: float) -> torch.Tensor: + device = v0.device if v0.dim() == 2: total_length = v0.shape[1] peak = int(total_length * (1 - t)) @@ -162,7 +163,7 @@ def merge_tensors_gradient(v0: torch.Tensor, v1: torch.Tensor, t: float) -> torc v0_ratios = 1 - blend_ratios # Vectorized blending of the tensors - result = (v1 * blend_ratios.unsqueeze(0)) + (v0 * v0_ratios.unsqueeze(0)) + result = (v1 * blend_ratios.unsqueeze(0).to(device)) + (v0 * v0_ratios.unsqueeze(0).to(device)) return result else: