Download model/cache.py from KiranAN1988/NeemSutra-125M-Python-Instruct-Beta: direct link, hf CLI and curl.
- Browser
- Download file 1.08 kB
-
https://huggingface.co/KiranAN1988/NeemSutra-125M-Python-Instruct-Beta/resolve/main/model/cache.py
- Command line
-
hf download hf://KiranAN1988/NeemSutra-125M-Python-Instruct-Beta/model/cache.py
-
curl -L -o cache.py https://huggingface.co/KiranAN1988/NeemSutra-125M-Python-Instruct-Beta/resolve/main/model/cache.py
1.08 kB
| import torch | |
| import torch.nn as nn | |
| class StaticKVCache(nn.Module): | |
| def __init__( | |
| self, | |
| max_batch_size, | |
| n_heads, | |
| context_length, | |
| head_dim, | |
| device, | |
| dtype=torch.float32 | |
| ): | |
| super().__init__() | |
| self.register_buffer( | |
| "k_cache", | |
| torch.zeros( | |
| (max_batch_size, n_heads, context_length, head_dim), | |
| dtype=dtype, | |
| device=device | |
| ) | |
| ) | |
| self.register_buffer( | |
| "v_cache", | |
| torch.zeros( | |
| (max_batch_size, n_heads, context_length, head_dim), | |
| dtype=dtype, | |
| device=device | |
| ) | |
| ) | |
| def update(self, k_val, v_val, cache_position): | |
| self.k_cache[:, :, cache_position] = k_val | |
| self.v_cache[:, :, cache_position] = v_val | |
| total_len = cache_position[-1].item() + 1 | |
| return ( | |
| self.k_cache[:, :, :total_len], | |
| self.v_cache[:, :, :total_len] | |
| ) |