From e417790413bd143dd52ed3847c339f964eb96d40 Mon Sep 17 00:00:00 2001 From: wassname Date: Mon, 13 Nov 2023 08:11:13 +0800 Subject: [PATCH] fix kv_cache faking? --- src/models/transformer.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/models/transformer.py b/src/models/transformer.py index ecf68ae..9b28849 100644 --- a/src/models/transformer.py +++ b/src/models/transformer.py @@ -130,10 +130,10 @@ class Transformer(nn.Module): # fake it, since it's used to keep track of steps if past_keys_values is not None: k_size = past_keys_values[0]._k_cache._cache.size() - k_size = (*k_size[:2], 1, *k_size[3:]) - v_size = past_keys_values[0]._v_cache._cache.size() - v_size = (*v_size[:2], 1, *v_size[3:]) - past_keys_values[0].update(torch.rand(k_size), torch.rand(v_size)) + # k_size = (x.shape[0], x.shape[1], x.shape[1], 1) + # v_size = past_keys_values[0]._v_cache._cache.size() + v_size = (k_size[0], k_size[1], x.shape[1], k_size[3]) + past_keys_values[0].update(torch.rand(v_size), torch.rand(v_size)) return x