From 9e84b150ec30e145abd3385c494bb392ef9338b7 Mon Sep 17 00:00:00 2001 From: Extraltodeus Date: Sat, 26 Apr 2025 01:10:16 +0200 Subject: [PATCH] slight quality improvement --- custom_samplers.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/custom_samplers.py b/custom_samplers.py index 059f585..aaf53cc 100644 --- a/custom_samplers.py +++ b/custom_samplers.py @@ -26,18 +26,16 @@ def fast_distance_weights(t, use_softmax=False, use_slerp=False, uncond=None): tn = t.div(norm) distances = (tn.unsqueeze(0) - tn.unsqueeze(1)).abs().sum(dim=0) + distances = distances.max(dim=0, keepdim=True).values - distances if uncond != None: uncond = uncond.div(torch.linalg.matrix_norm(uncond, keepdim=True)) - distances -= tn.sub(uncond).abs() + distances += tn.sub(uncond).abs().div(n) if use_softmax: - distances = distances.max(dim=0).values - distances distances = distances.mul(n).softmax(dim=0) else: - distances = 1 - (distances - distances.min(dim=0).values) / (distances.max(dim=0).values - distances.min(dim=0).values) - distances[~torch.isfinite(distances)] = 1 - distances = distances.pow(2) + distances = distances.div(distances.max(dim=0).values).pow(2) distances = distances / distances.sum(dim=0) if use_slerp: