compressor-reflex / serializer.py
aialchemist-dev's picture
Compressor Reflex v2 (INT8): 100% must-keep retention @ tau*=0.50
11d7a2c verified
Raw History Blame Contribute Delete
5 kB
import os
from typing import List, Tuple, Dict, Any
class CompressorSerializer:
"""
Serializes tool outputs and step intents into tokenized chunks with line spans.
- Intent capped at 128 tokens
- Chunk capped at 512 tokens
- Sliding window with line preservation
"""
def __init__(self, tokenizer, max_intent_tokens: int = 128, max_chunk_tokens: int = 512):
self.tokenizer = tokenizer
self.max_intent_tokens = max_intent_tokens
self.max_chunk_tokens = max_chunk_tokens
def chunk_tool_output(self, intent: str, raw_output: str) -> List[Dict[str, Any]]:
"""
Splits raw_output into chunks <= max_chunk_tokens, preserving whole lines.
Returns a list of chunks, each with:
- chunk_text
- lines (list of str)
- input_ids, attention_mask
- line_token_spans: list of (start_token, end_token)
"""
raw_lines = raw_output.splitlines()
if not raw_lines:
return []
# Tokenize intent with truncation
intent_enc = self.tokenizer(
f"INTENT: {intent.strip()}\n---\n",
add_special_tokens=False,
truncation=True,
max_length=self.max_intent_tokens
)
intent_prefix = self.tokenizer.decode(intent_enc["input_ids"])
intent_token_len = len(intent_enc["input_ids"])
avail_chunk_tokens = self.max_chunk_tokens - intent_token_len - 4 # room for CLS/SEP
chunks = []
curr_lines = []
curr_tokens_est = 0
for line in raw_lines:
line_str = line if line.strip() else " "
# Fast token estimation: ~3.5 chars per token
line_tok_est = max(1, len(line_str) // 3 + 1)
if curr_lines and (curr_tokens_est + line_tok_est > avail_chunk_tokens):
# Finalize current chunk
chunk_obj = self._build_chunk(intent_prefix, curr_lines)
chunks.append(chunk_obj)
curr_lines = [line_str]
curr_tokens_est = line_tok_est
else:
curr_lines.append(line_str)
curr_tokens_est += line_tok_est
if curr_lines:
chunk_obj = self._build_chunk(intent_prefix, curr_lines)
chunks.append(chunk_obj)
return chunks
def _build_chunk(self, intent_prefix: str, lines: List[str]) -> Dict[str, Any]:
"""
Formats chunk text and computes line token spans via offset mapping.
"""
prefix = intent_prefix
chunk_body = "\n".join(lines)
full_text = prefix + chunk_body
enc = self.tokenizer(
full_text,
return_offsets_mapping=True,
truncation=True,
max_length=self.max_chunk_tokens,
return_tensors="pt"
)
offsets = enc["offset_mapping"][0].tolist()
line_spans = []
curr_char = len(prefix)
for line in lines:
start_char = curr_char
end_char = curr_char + len(line)
curr_char = end_char + 1 # newline
t_start = None
t_end = None
for idx, (os_char, oe_char) in enumerate(offsets):
if os_char == oe_char:
continue
if os_char >= start_char and t_start is None:
t_start = idx
if oe_char <= end_char and t_start is not None:
t_end = idx + 1
if t_start is not None and t_end is not None and t_end > t_start:
line_spans.append((t_start, t_end))
elif t_start is not None:
line_spans.append((t_start, min(t_start + 1, len(offsets) - 1)))
else:
# Truncated or zero-token line fallback
last_idx = max(0, len(offsets) - 2)
line_spans.append((last_idx, last_idx + 1))
return {
"full_text": full_text,
"lines": lines,
"input_ids": enc["input_ids"][0],
"attention_mask": enc["attention_mask"][0],
"line_spans": line_spans,
"num_lines": len(lines)
}
if __name__ == "__main__":
from transformers import AutoTokenizer
tok = AutoTokenizer.from_pretrained(r"C:\Users\EricM\Dev\jev-research\compressor-reflex\base_model")
serializer = CompressorSerializer(tok)
intent = "Inspect pytest output for failure in auth"
output = "running pytest...\n==== test session starts ====\nFAILED test_auth.py:23 - Token expired\n==== 1 failed ===="
chunks = serializer.chunk_tool_output(intent, output)
print(f"Generated {len(chunks)} chunk(s)")
for i, c in enumerate(chunks):
print(f"Chunk {i}: {c['num_lines']} lines, {len(c['input_ids'])} tokens")
for l_idx, (s, e) in enumerate(c["line_spans"]):
line_text = c["lines"][l_idx]
print(f" Line {l_idx} [{s}:{e}]: {line_text}")