Quazim0t0 commited on
Commit
42194bf
·
1 Parent(s): bfca7ea

Harness: pass engram_context_ids on cached steps (n-gram memory matches full compute; exercised in real-weight A/B) + no-repeat-ngram(3)

Browse files
Files changed (1) hide show
  1. generate_triattention.py +8 -1
generate_triattention.py CHANGED
@@ -77,6 +77,12 @@ def generate(model, tok, prompt, max_new_tokens, temperature, top_k, device,
77
  t0 = time.time()
78
  for _step in range(max_new_tokens):
79
  logits = _rep_penalty(logits, generated[0, P:], repetition_penalty)
 
 
 
 
 
 
80
  if temperature and temperature > 0:
81
  probs = torch.softmax(logits / temperature, dim=-1)
82
  if top_k:
@@ -91,8 +97,9 @@ def generate(model, tok, prompt, max_new_tokens, temperature, top_k, device,
91
  if tok.eos_token_id is not None and nxt.item() == tok.eos_token_id:
92
  break
93
  cur_pos = torch.tensor([[P + _step]], device=device)
 
94
  out = model(input_ids=nxt, position_ids=cur_pos,
95
- past_key_values=pkv, use_cache=True)
96
  pkv = out.past_key_values
97
  logits = out.logits[:, -1, :]
98
  cache_len_series.append(pkv[0][0].shape[2])
 
77
  t0 = time.time()
78
  for _step in range(max_new_tokens):
79
  logits = _rep_penalty(logits, generated[0, P:], repetition_penalty)
80
+ gen = generated[0, P:].tolist()
81
+ if len(gen) >= 3: # no-repeat-ngram(3): A/B'd on real
82
+ t2 = (gen[-2], gen[-1]) # weights, rep-4 -> 0.000
83
+ for i in range(len(gen) - 2):
84
+ if (gen[i], gen[i + 1]) == t2:
85
+ logits[0, gen[i + 2]] = float("-inf")
86
  if temperature and temperature > 0:
87
  probs = torch.softmax(logits / temperature, dim=-1)
88
  if top_k:
 
97
  if tok.eos_token_id is not None and nxt.item() == tok.eos_token_id:
98
  break
99
  cur_pos = torch.tensor([[P + _step]], device=device)
100
+ ctx = generated[:, -3:-1] # engram trigram context (2 tokens before nxt)
101
  out = model(input_ids=nxt, position_ids=cur_pos,
102
+ past_key_values=pkv, use_cache=True, engram_context_ids=ctx)
103
  pkv = out.past_key_values
104
  logits = out.logits[:, -1, :]
105
  cache_len_series.append(pkv[0][0].shape[2])