Instructions to use cleopatro/context_tracking with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use cleopatro/context_tracking with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("token-classification", model="cleopatro/context_tracking")# pip install -U transformers accelerate # Load model directly from transformers import AutoTokenizer, AutoModelForTokenClassification tokenizer = AutoTokenizer.from_pretrained("cleopatro/context_tracking") model = AutoModelForTokenClassification.from_pretrained("cleopatro/context_tracking", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download stm.py from cleopatro/context_tracking: direct link, hf CLI and curl.
- Browser
- Download file 1.84 kB
-
https://huggingface.co/cleopatro/context_tracking/resolve/main/stm.py
- Command line
-
hf download hf://cleopatro/context_tracking/stm.py
-
curl -L -o stm.py https://huggingface.co/cleopatro/context_tracking/resolve/main/stm.py
1.84 kB
| from collections import defaultdict | |
| class ShortTermMemory: | |
| def __init__(self, window_size=10, decay_rate=0.5): | |
| self.abstract_entities = defaultdict(int) | |
| self.locations = defaultdict(int) | |
| self.times = defaultdict(int) | |
| self.window_size = window_size | |
| self.decay_rate = decay_rate | |
| def update(self, entity_type, entity): | |
| # Determine the appropriate dictionary based on the entity type | |
| if entity_type == 'abstract': | |
| entity_dict = self.abstract_entities | |
| elif entity_type == 'location': | |
| entity_dict = self.locations | |
| elif entity_type == 'time': | |
| entity_dict = self.times | |
| else: | |
| raise ValueError(f'Invalid entity type: {entity_type}') | |
| # Increment the count for the given entity | |
| entity_dict[entity] += 1 | |
| # Decay the counts of other entities in the same dictionary | |
| for e, count in list(entity_dict.items()): | |
| if e != entity: | |
| entity_dict[e] = int(count * self.decay_rate) | |
| # Remove entities with count <= 1 | |
| entity_dict = {e: count for e, count in entity_dict.items() if count > 1} | |
| # Trim the dictionary to the window size | |
| entity_dict = dict(sorted(entity_dict.items(), key=lambda x: x[1], reverse=True)[:self.window_size]) | |
| # Update the appropriate dictionary with the trimmed version | |
| if entity_type == 'abstract': | |
| self.abstract_entities = entity_dict | |
| elif entity_type == 'location': | |
| self.locations = entity_dict | |
| elif entity_type == 'time': | |
| self.times = entity_dict | |
| def get_memory(self): | |
| return { | |
| 'abstract_entities': self.abstract_entities, | |
| 'locations': self.locations, | |
| 'times': self.times | |
| } | |