| 43 | |
| 44 | class KVCacheTest(unittest.TestCase): |
| 45 | def setUp(self): |
| 46 | self.batch_size = 1 |
| 47 | self.max_seq_len = 10 |
| 48 | self.num_kv_heads = 1 # For testing purposes, usually this is 64. |
| 49 | self.head_dim = 8 |
| 50 | self.dtype = torch.float |
| 51 | |
| 52 | self.tt_kv_cache = KVCache( |
| 53 | batch_size=self.batch_size, |
| 54 | max_seq_len=self.max_seq_len, |
| 55 | num_kv_heads=self.num_kv_heads, |
| 56 | head_dim=self.head_dim, |
| 57 | dtype=self.dtype, |
| 58 | ) |
| 59 | self.et_kv_cache = InferenceKVCache( |
| 60 | batch_size=self.batch_size, |
| 61 | max_seq_len=self.max_seq_len, |
| 62 | num_kv_heads=self.num_kv_heads, |
| 63 | head_dim=self.head_dim, |
| 64 | dtype=self.dtype, |
| 65 | transpose_cache=False, |
| 66 | ) |
| 67 | |
| 68 | def _test_kv_cache(self, et_cache_module: Callable): |
| 69 | """ |