Instructions to use Synthyra/Profluent-E1-150M with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Synthyra/Profluent-E1-150M with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("fill-mask", model="Synthyra/Profluent-E1-150M", trust_remote_code=True)# Load model directly from transformers import AutoModelForMaskedLM model = AutoModelForMaskedLM.from_pretrained("Synthyra/Profluent-E1-150M", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """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]) | |