nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
7.83 kB
"""Typed CKY-style parser over locally scored semantic fragments."""
from __future__ import annotations
from dataclasses import dataclass, replace
import math
from strata.modeling.compose.ast import ProgramNode, apply_chain, canonicalize
from strata.modeling.compose.fragments import FunctionFragment, LexicalCandidate
from strata.modeling.compose.types import SemanticType
@dataclass(frozen=True, slots=True)
class ChartItem:
start: int
end: int
score: float
inside_score: float
function: FunctionFragment | None = None
program: ProgramNode | None = None
derivation: tuple[str, ...] = ()
def __post_init__(self) -> None:
if (self.function is None) == (self.program is None):
raise ValueError("chart item must contain exactly one function or value program")
@property
def output_type(self) -> SemanticType:
if self.function is not None:
return self.function.output_type
assert self.program is not None
return self.program.output_type
def key(self) -> tuple:
if self.program is not None:
return ("value", canonicalize(self.program))
assert self.function is not None
return (
"function",
self.function.argument_side,
tuple(operator.value for operator in self.function.operators),
)
@dataclass(frozen=True, slots=True)
class ParseResult:
program: ProgramNode
score: float
inside_score: float
derivation: tuple[str, ...]
class TypedChartParser:
"""Compose only type-compatible fragments; invalid ASTs never enter the chart."""
def __init__(self, *, beam_size: int = 32, combination_score: float = 0.0) -> None:
if beam_size <= 0:
raise ValueError("beam_size must be positive")
self.beam_size = int(beam_size)
self.combination_score = float(combination_score)
def parse(
self,
length: int,
candidates: tuple[LexicalCandidate, ...],
*,
required_output: SemanticType | None = None,
) -> ParseResult:
if length <= 0:
raise ValueError("length must be positive")
chart: dict[tuple[int, int], dict[tuple, ChartItem]] = {}
for candidate in candidates:
if candidate.end > length:
raise ValueError("lexical candidate lies outside the chart")
item = ChartItem(
candidate.start,
candidate.end,
candidate.score,
candidate.score,
function=candidate.function,
program=candidate.anchor,
derivation=(f"lex:{candidate.start}:{candidate.end}",),
)
self._insert(chart.setdefault((item.start, item.end), {}), item)
for width in range(2, length + 1):
for start in range(0, length - width + 1):
end = start + width
cell = chart.setdefault((start, end), {})
for split in range(start + 1, end):
left = chart.get((start, split), {})
right = chart.get((split, end), {})
for left_item in left.values():
for right_item in right.values():
for combined in self._combine(left_item, right_item):
self._insert(cell, combined)
if len(cell) > self.beam_size:
kept = sorted(
cell.values(),
key=lambda item: (-item.score, repr(item.key())),
)[: self.beam_size]
chart[(start, end)] = {item.key(): item for item in kept}
complete = [item for item in chart.get((0, length), {}).values() if item.program is not None]
if required_output is not None:
complete = [item for item in complete if item.output_type is required_output]
if not complete:
raise ValueError("no complete well-typed derivation")
best = sorted(complete, key=lambda item: (-item.score, repr(item.key())))[0]
assert best.program is not None
return ParseResult(best.program, best.score, best.inside_score, best.derivation)
def _combine(self, left: ChartItem, right: ChartItem) -> tuple[ChartItem, ...]:
score = left.score + right.score + self.combination_score
inside = left.inside_score + right.inside_score + self.combination_score
derivation = left.derivation + right.derivation + (f"combine:{left.end}",)
outputs: list[ChartItem] = []
if left.program is not None and right.function is not None:
if right.function.argument_side == "left" and left.output_type is right.function.input_type:
outputs.append(ChartItem(
left.start,
right.end,
score,
inside,
program=apply_chain(left.program, right.function.operators),
derivation=derivation,
))
if left.function is not None and right.program is not None:
if left.function.argument_side == "right" and right.output_type is left.function.input_type:
outputs.append(ChartItem(
left.start,
right.end,
score,
inside,
program=apply_chain(right.program, left.function.operators),
derivation=derivation,
))
if left.function is not None and right.function is not None:
if (
left.function.argument_side == right.function.argument_side == "left"
and left.function.output_type is right.function.input_type
):
outputs.append(ChartItem(
left.start,
right.end,
score,
inside,
function=FunctionFragment(
left.function.operators + right.function.operators,
argument_side="left",
),
derivation=derivation,
))
if (
left.function.argument_side == right.function.argument_side == "right"
and right.function.output_type is left.function.input_type
):
outputs.append(ChartItem(
left.start,
right.end,
score,
inside,
function=FunctionFragment(
right.function.operators + left.function.operators,
argument_side="right",
),
derivation=derivation,
))
return tuple(outputs)
@staticmethod
def _insert(cell: dict[tuple, ChartItem], candidate: ChartItem) -> None:
key = candidate.key()
previous = cell.get(key)
if previous is None:
cell[key] = candidate
return
inside = float(torch_logaddexp(previous.inside_score, candidate.inside_score))
if candidate.score > previous.score:
winner = candidate
elif candidate.score < previous.score:
winner = previous
else:
winner = min((candidate, previous), key=lambda item: repr(item.derivation))
cell[key] = replace(winner, inside_score=inside)
def torch_logaddexp(left: float, right: float) -> float:
maximum = max(left, right)
if math.isinf(maximum) and maximum < 0:
return maximum
return maximum + math.log(math.exp(left - maximum) + math.exp(right - maximum))
__all__ = ["ChartItem", "ParseResult", "TypedChartParser"]