File size: 5,030 Bytes
82f262a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
"""
RegexFeatureExtractor - token feature gating.
Extends build_code_features (4-dim) to a 7-dim per-token feature vector plus
a running SyntaxState (bracket stack, quote open/closed) used as a cheap
pre-model sanity gate.

Dimensions per token:
  [0] is_code_like       (indent / brackets / operators / newlines)
  [1] indent_depth       (normalized leading whitespace)
  [2] bracket_balance    (1.0 open, 0.5 neutral, 0.0 close)
  [3] has_newline
  [4] keyword_hit        (NEW)
  [5] quote_state        (running open/closed string literal)   (NEW)
  [6] numeric_literal    (NEW)
"""

import re
import torch
from dataclasses import dataclass, field
from typing import List, Optional, Tuple

CODE_CHARS = set("{}()[];=<>!&|+-*/%'\"`#@.,:")
KEYWORD_RE = re.compile(
    r"\b(def|class|import|return|if|else|elif|for|while|try|except|finally|"
    r"lambda|pass|with|as|yield|from|async|await)\b"
)
NUMERIC_RE = re.compile(r"\b\d+(\.\d+)?\b")

OPEN_BRACKETS = {"{": 1, "[": 1, "(": 1}
CLOSE_BRACKETS = {"}": 1, "]": 1, ")": 1}


@dataclass
class SyntaxState:
    bracket_stack: List[str] = field(default_factory=list)
    quote_open: Optional[str] = None  # '"' or "'" while a string literal is open
    depth: int = 0

    def is_balanced(self) -> bool:
        return not self.bracket_stack and self.quote_open is None

    def to_dict(self) -> dict:
        return {
            "balanced": self.is_balanced(),
            "bracket_depth": len(self.bracket_stack),
            "quote_open": self.quote_open,
        }


class RegexFeatureExtractor:
    def __init__(self, num_features: int = 7):
        self.num_features = num_features

    def extract(self, tokenizer, input_ids: torch.Tensor) -> Tuple[torch.Tensor, List[SyntaxState]]:
        """Per-token features (B, T, F) + one SyntaxState per row."""
        feats: List[List[List[float]]] = []
        states: List[SyntaxState] = []
        for row in input_ids.tolist():
            tokens = tokenizer.convert_ids_to_tokens(row)
            state = SyntaxState()
            row_feats = []
            for tok in tokens:
                is_code = any(c in CODE_CHARS for c in tok)
                indent = 0.0
                stripped = tok.lstrip()
                if stripped and tok != stripped:
                    indent = min((len(tok) - len(stripped)) / 8.0, 1.0)
                    is_code = True
                bal = 0.5
                for c in tok:
                    if c in OPEN_BRACKETS:
                        bal = 1.0
                        state.bracket_stack.append(c)
                    elif c in CLOSE_BRACKETS:
                        bal = 0.0
                        if state.bracket_stack:
                            state.bracket_stack.pop()
                # running quote state
                for c in tok:
                    if c in ('"', "'"):
                        if state.quote_open is None:
                            state.quote_open = c
                        elif state.quote_open == c:
                            state.quote_open = None
                quote = 1.0 if state.quote_open is not None else 0.0
                newline = 1.0 if "\n" in tok else 0.0
                kw = 1.0 if KEYWORD_RE.search(tok) else 0.0
                num = 1.0 if NUMERIC_RE.search(tok) else 0.0
                row_feats.append([1.0 if is_code else 0.0, indent, bal, newline, kw, quote, num])
                state.depth = len(state.bracket_stack)
            row_feats = row_feats[: input_ids.shape[1]]
            feats.append(row_feats)
            states.append(state)
        max_len = max(len(r) for r in feats)
        padded = [
            r + [[0.0, 0.0, 0.5, 0.0, 0.0, 0.0, 0.0]] * (max_len - len(r))
            for r in feats
        ]
        t = torch.tensor(padded, dtype=torch.float32)
        if t.shape[-1] > self.num_features:
            t = t[..., : self.num_features]
        return t, states

    def extract_text(self, text: str) -> SyntaxState:
        """Run a string-only pass for the pre-model gate (no tokenizer)."""
        state = SyntaxState()
        for c in text:
            if c in OPEN_BRACKETS:
                state.bracket_stack.append(c)
            elif c in CLOSE_BRACKETS and state.bracket_stack:
                state.bracket_stack.pop()
            elif c in ('"', "'"):
                if state.quote_open is None:
                    state.quote_open = c
                elif state.quote_open == c:
                    state.quote_open = None
        state.depth = len(state.bracket_stack)
        return state

    def gate(self, syntax_state: SyntaxState) -> str:
        """Returns 'pass' | 'warn' | 'block'."""
        if syntax_state.is_balanced():
            return "pass"
        return "block" if syntax_state.depth > 4 else "warn"


# drop-in replacement for architecture.build_code_features with 7-dim output
def build_code_features_v2(tokenizer, input_ids: torch.Tensor) -> torch.Tensor:
    return RegexFeatureExtractor(num_features=7).extract(tokenizer, input_ids)[0]