Views
No views yet
Llama-3-8B-Instruct packaged with custom code for the InferenceMemoryWrapper.trust_remote_code=True.InferenceMemoryWrapper1from transformers import AutoModelForCausalLM, AutoTokenizer
2import torch # Added import for example
3
4model_id = "your-username/your-repo-name" # Replace with your repo ID
5
6# Load the model and tokenizer, allowing custom code execution
7# Requires sufficient VRAM for the Llama 8B model + memory buffer
8model = AutoModelForCausalLM.from_pretrained(
9 model_id,
10 trust_remote_code=True,
11 torch_dtype=torch.float16, # Recommended for memory
12 device_map="auto"
13)
14tokenizer = AutoTokenizer.from_pretrained(model_id)
15
16# Example prompt
17prompt = "What is the capital of France?"
18inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
19
20# Generate using the custom method
21# Note: The memory buffer is initially randomly initialized unless loaded separately.
22# It will be updated during generation if update_rule is 'ema' or 'surprise'.
23outputs = model.generate(
24 **inputs,
25 max_new_tokens=50,
26 use_memory=True,
27 update_rule='ema' # or 'surprise' or 'none'
28)
29
30generated_text = tokenizer.decode(outputs[0], skip_special_tokens=True)
31print(generated_text)
32
33# To save user-specific memory state (after generation/updates):
34# user_memory_state = model.memory_buffer.data.clone()
35# user_surprise_state = model.surprise_state.clone()
36# torch.save({'memory_buffer': user_memory_state, 'surprise_state': user_surprise_state}, 'user_memory.pt')
37
38# To load user-specific memory state:
39# loaded_state = torch.load('user_memory.pt')
40# model.memory_buffer.data.copy_(loaded_state['memory_buffer'])
41# model.surprise_state.copy_(loaded_state['surprise_state'])
42memory_buffer and surprise_state in this packaged model are initialized randomly according to the InferenceMemoryWrapper code. They do not contain any pre-trained memory state unless you load it separately after initializing the model (see example above). You need to manage loading/saving the memory state per user externally.