Instructions to use Synthyra/Profluent-E1-300M with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Synthyra/Profluent-E1-300M with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("fill-mask", model="Synthyra/Profluent-E1-300M", trust_remote_code=True)# Load model directly from transformers import AutoModelForMaskedLM model = AutoModelForMaskedLM.from_pretrained("Synthyra/Profluent-E1-300M", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 9,987 Bytes
77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 b3e55d2 77372f3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 | """Key-value cache implementations used by E1 inference."""
from __future__ import annotations
import torch
from typing import Any
from transformers.modeling_outputs import ModelOutput
from transformers.utils import logging
def _get_logger():
"""Resolve the Transformers logger only when a cache path emits a message."""
return logging.get_logger(__name__)
class DynamicCache:
"""A cache that grows K and V along their sequence dimension.
Each cached tensor has shape (b, l, h, d).
Args:
key_cache (`list[torch.Tensor]`): The list of key states.
value_cache (`list[torch.Tensor]`): The list of value states.
"""
def __init__(self) -> None:
self.key_cache: list[torch.Tensor] = []
self.value_cache: list[torch.Tensor] = []
def update(
self, key_states: torch.Tensor, value_states: torch.Tensor, layer_idx: int
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Update the key and value caches in-place, and return the necessary keys and value states.
Args:
key_states (`torch.Tensor`): K to cache with shape (b, l, h, d).
value_states (`torch.Tensor`): V to cache with shape (b, l, h, d).
layer_idx (`int`): The index of the layer to update.
Returns:
tuple[`torch.Tensor`, `torch.Tensor`]: Cached K and V, each with shape
(b, l, h, d).
"""
# key_states, value_states: (b, l_new, h, d)
if len(self.key_cache) <= layer_idx:
# Empty tensors preserve skipped layer indices until those layers receive state.
for _ in range(len(self.key_cache), layer_idx):
self.key_cache.append(torch.tensor([]))
self.value_cache.append(torch.tensor([]))
self.key_cache.append(key_states)
self.value_cache.append(value_states)
elif (
not self.key_cache[
layer_idx
].numel() # prefers not t.numel() to len(t) == 0 to export the model
): # fills previously skipped layers; checking for tensor causes errors
self.key_cache[layer_idx] = key_states
self.value_cache[layer_idx] = value_states
else:
self.key_cache[layer_idx] = torch.cat( # (b, l_cached + l_new, h, d)
[self.key_cache[layer_idx], key_states],
dim=1,
)
self.value_cache[layer_idx] = torch.cat( # (b, l_cached + l_new, h, d)
[self.value_cache[layer_idx], value_states], dim=1
)
return ( # (b, l_total, h, d), (b, l_total, h, d)
self.key_cache[layer_idx],
self.value_cache[layer_idx],
)
def get_seq_length(self, layer_idx: int = 0) -> int:
"""Return the cached sequence length for one layer."""
is_empty_layer = (
len(self.key_cache) == 0 # no cache in any layer
or len(self.key_cache)
<= layer_idx # skipped `layer_idx` and hasn't run a layer with cache after it
or not self.key_cache[layer_idx].numel() # the layer has no cache
)
layer_seq_length = self.key_cache[layer_idx].shape[1] if not is_empty_layer else 0
return layer_seq_length
def crop(self, max_length: int) -> None:
"""Crop every cached K and V tensor to ``max_length`` tokens."""
if max_length <= 0:
raise ValueError("max_length must be positive")
if self.get_seq_length() <= max_length:
return
for layer_idx in range(len(self.key_cache)):
if self.key_cache[layer_idx].numel():
self.key_cache[layer_idx] = self.key_cache[layer_idx][:, :max_length, ...]
self.value_cache[layer_idx] = self.value_cache[layer_idx][:, :max_length, ...]
def batch_repeat_interleave(self, repeats: int) -> None:
"""Repeat the cache `repeats` times in the batch dimension. Used in contrastive search."""
for layer_idx in range(len(self.key_cache)):
if self.key_cache[layer_idx].numel():
# (b * repeats, l, h, d)
self.key_cache[layer_idx] = self.key_cache[layer_idx].repeat_interleave(
repeats, dim=0
)
self.value_cache[layer_idx] = self.value_cache[layer_idx].repeat_interleave(
repeats, dim=0
) # (b * repeats, l, h, d)
def batch_select_indices(self, indices: torch.Tensor | list[int]) -> None:
"""Keep selected rows of the cache batch dimension."""
for layer_idx in range(len(self.key_cache)):
if self.key_cache[layer_idx].numel():
self.key_cache[layer_idx] = self.key_cache[layer_idx][indices, ...] # (n, l, h, d)
self.value_cache[layer_idx] = self.value_cache[layer_idx][
indices, ...
] # (n, l, h, d)
class KVCache:
def __init__(self, cache_size: int = 4) -> None:
self.cache_size = cache_size
self.tensor_input_field_names = [
"input_ids",
"within_seq_position_ids",
"global_position_ids",
"sequence_ids",
"labels",
]
# Upstream E1 called the encoder output ``embeddings``. FastPLMs uses
# the standard Transformers ``last_hidden_state`` name, while keeping
# the aliases here makes the cache safe for either output contract.
self.tensor_output_field_names = [
"logits",
"last_hidden_state",
"embeddings",
"token_embeddings",
]
self.cache_dict: dict[str, DynamicCache] = {}
self.cache_queue: list[str] = []
def reset(self) -> None:
for k in list(self.cache_dict.keys()):
del self.cache_dict[k]
del self.cache_dict
self.cache_dict = {}
self.cache_queue = []
torch.cuda.empty_cache()
def before_forward(self, batch: dict[str, torch.Tensor]) -> None:
contexts: list[str] | None = batch.get("context")
if contexts is None or "context_len" not in batch:
_get_logger().warning_once(
"KVCache requires both `context` and `context_len`; cache setup was skipped."
)
return
context_lens: list[int] = list(set(batch["context_len"]))
contexts: list[str] = list(set(contexts)) # type: ignore[no-redef]
if len(contexts) != 1 or len(context_lens) != 1:
_get_logger().warning(
"SingleContextKVCache requires a single context and context length. "
"Multiple contexts or context lengths found in a single batch. Skipping."
)
return
batch_size = batch["input_ids"].shape[0] # b
unique_context = contexts[0]
unique_context_len = context_lens[0]
batch["use_cache"] = True
if unique_context not in self.cache_dict:
return
self.cache_dict[unique_context].batch_repeat_interleave(batch_size)
past_key_values = self.cache_dict[unique_context]
batch["past_key_values"] = past_key_values
# A cached prefix leaves only query-suffix tokens for the model call.
for field_name in self.tensor_input_field_names:
if batch.get(field_name) is not None:
batch[field_name] = batch[field_name][:, unique_context_len:] # (b, l_suffix, ...)
def after_forward(self, batch: dict[str, Any], outputs: ModelOutput) -> None:
contexts = batch.get("context")
context_lens = batch.get("context_len", [])
if (
contexts is None
or len(set(contexts)) != 1
or len(set(context_lens)) != 1
or context_lens[0] == 0
):
return
if not batch.get("use_cache", False):
raise ValueError("E1 retrieval cache updates require use_cache=True.")
unique_context = contexts[0]
unique_context_len = context_lens[0]
past_key_values = getattr(outputs, "past_key_values", None)
if not isinstance(past_key_values, DynamicCache):
_get_logger().warning_once(
"KVCache is incompatible with models that don't return a DynamicCache. Skipping."
)
return
if "past_key_values" not in batch:
if len(self.cache_queue) == self.cache_size:
last_context = self.cache_queue.pop(0)
if last_context not in self.cache_queue:
del self.cache_dict[last_context]
torch.cuda.empty_cache()
self.cache_dict[unique_context] = past_key_values
self.cache_queue.append(unique_context)
# The first uncached call returns the full sequence; expose its query suffix.
for field_name in self.tensor_input_field_names:
if field_name in batch and batch[field_name] is not None:
batch[field_name] = batch[field_name][
:, unique_context_len:
] # (b, l_suffix, ...)
for field_name in self.tensor_output_field_names:
if field_name in outputs and outputs[field_name] is not None:
outputs[field_name] = outputs[field_name][
:, unique_context_len:
] # (b, l_suffix, ...)
if "hidden_states" in outputs and outputs["hidden_states"] is not None:
hidden_states = outputs["hidden_states"]
sliced_hidden_states = tuple( # each: (b, l_suffix, d)
hidden_state[:, unique_context_len:] for hidden_state in hidden_states
)
outputs["hidden_states"] = sliced_hidden_states
self.cache_dict[unique_context].crop(unique_context_len)
self.cache_dict[unique_context].batch_select_indices([0])
|