autodiff-visualizer / input_parser.py
Simon Clematide
Add compact vector and matrix support with adjoint backpropagation
6a7acce
Raw History Blame Contribute Delete
5.68 kB
import ast
import re
import numpy as np
MAX_INPUT_TEXT_LEN = 1000
MAX_TOTAL_ELEMENTS = 50000
MAX_ARRAY_RANK = 2
MAX_NUMERIC_MAGNITUDE = 1e9
def split_top_level_assignments(text):
"""
Split an assignment string (e.g. 'W = [[1, 2], [3, 4]], x = [1, 1], b = 0.5')
only at commas that are at bracket depth 0.
"""
chunks = []
current = []
depth = 0
for char in text:
if char in '([{':
depth += 1
current.append(char)
elif char in ')]}':
depth -= 1
if depth < 0:
raise ValueError("Unmatched closing bracket in variable values.")
current.append(char)
elif char == ',' and depth == 0:
chunks.append(''.join(current).strip())
current = []
else:
current.append(char)
if depth != 0:
raise ValueError("Unmatched opening bracket in variable values.")
last = ''.join(current).strip()
if last:
chunks.append(last)
return chunks
def _validate_and_convert_literal(node):
"""
Recursively inspect an AST literal node and convert to a Python float or nested list of floats.
Strictly forbids function calls, comprehensions, variables, strings, booleans, etc.
"""
if isinstance(node, ast.Constant):
# In Python 3.8+, booleans are an instance of int/Constant with value bool
if isinstance(node.value, bool):
raise ValueError("Boolean literals are not allowed.")
if not isinstance(node.value, (int, float)):
raise ValueError(f"Only numeric literals are allowed, got {type(node.value).__name__}.")
val = float(node.value)
if np.isnan(val) or np.isinf(val):
raise ValueError("Values must be finite real numbers.")
if abs(val) > MAX_NUMERIC_MAGNITUDE:
raise ValueError(f"Numeric literal exceeds maximum allowed magnitude ({MAX_NUMERIC_MAGNITUDE}).")
return val
if hasattr(ast, "Num") and isinstance(node, ast.Num):
val = float(node.n)
if abs(val) > MAX_NUMERIC_MAGNITUDE:
raise ValueError(f"Numeric literal exceeds maximum allowed magnitude.")
return val
if isinstance(node, ast.UnaryOp):
if isinstance(node.op, ast.UAdd):
return +_validate_and_convert_literal(node.operand)
elif isinstance(node.op, ast.USub):
return -_validate_and_convert_literal(node.operand)
raise ValueError(f"Unsupported unary operator in literal: {type(node.op).__name__}")
if isinstance(node, ast.List):
if len(node.elts) == 0:
raise ValueError("Empty lists or arrays are not allowed.")
return [_validate_and_convert_literal(elt) for elt in node.elts]
raise ValueError(f"Disallowed construct in literal: {type(node).__name__}")
def parse_literal_value(val_str):
"""
Safely parse a literal string (number, vector list, or matrix list) into a float or np.ndarray.
Enforces rectangular matrix shape, finite elements, rank <= 2, and element count budget.
"""
val_str = val_str.strip()
if not val_str:
raise ValueError("Empty value literal.")
try:
parsed = ast.parse(val_str, mode="eval")
except SyntaxError as e:
raise ValueError(f"Syntax error in literal '{val_str}': {e.msg}")
val = _validate_and_convert_literal(parsed.body)
if isinstance(val, (int, float)):
return float(val)
# Convert nested lists to numpy array
try:
arr = np.array(val, dtype=np.float64)
except ValueError as e:
raise ValueError(f"Malformed or ragged array structure: {e}")
if arr.ndim > MAX_ARRAY_RANK:
raise ValueError(f"Array rank {arr.ndim} exceeds maximum supported rank ({MAX_ARRAY_RANK}).")
if arr.size == 0:
raise ValueError("Empty array dimensions are not allowed.")
if arr.size > MAX_TOTAL_ELEMENTS:
raise ValueError(f"Array size ({arr.size} elements) exceeds limit ({MAX_TOTAL_ELEMENTS}).")
if not np.all(np.isfinite(arr)):
raise ValueError("Array contains nonfinite numbers (NaN or Inf).")
return arr
def parse_variable_assignments(text):
"""
Parse a comma-separated string of variable assignments into a dictionary:
{ "W": np.ndarray, "x": np.ndarray, "b": float, ... }
Validates variable names, duplicate keys, and bounds.
"""
if text is None or not text.strip():
return {}
cleaned = text.strip()
if len(cleaned) > MAX_INPUT_TEXT_LEN:
raise ValueError(f"Input text exceeds maximum length ({MAX_INPUT_TEXT_LEN} characters).")
chunks = split_top_level_assignments(cleaned)
feed_dict = {}
total_elements = 0
for chunk in chunks:
if not chunk:
continue
if "=" not in chunk:
raise ValueError(f"Expected assignment 'var = value', got '{chunk}'")
var_part, val_part = chunk.split("=", 1)
var_name = var_part.strip()
val_str = val_part.strip()
if not re.match(r"^[a-zA-Z][a-zA-Z0-9_]*$", var_name):
raise ValueError(f"Invalid variable identifier: '{var_name}'")
if var_name in feed_dict:
raise ValueError(f"Duplicate assignment for variable '{var_name}'")
parsed_val = parse_literal_value(val_str)
if isinstance(parsed_val, np.ndarray):
total_elements += parsed_val.size
else:
total_elements += 1
if total_elements > MAX_TOTAL_ELEMENTS:
raise ValueError(f"Total element count exceeds limit ({MAX_TOTAL_ELEMENTS}).")
feed_dict[var_name] = parsed_val
return feed_dict