calc-flow-python 4.0.0__cp313-abi3-win_amd64.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- calc_flow/__init__.py +169 -0
- calc_flow/_native.pyd +0 -0
- calc_flow/_native.pyi +366 -0
- calc_flow/array.py +1324 -0
- calc_flow/capabilities.py +775 -0
- calc_flow/config.py +219 -0
- calc_flow/errors.py +25 -0
- calc_flow/join_spec.py +156 -0
- calc_flow/pipeline.py +1123 -0
- calc_flow/py.typed +0 -0
- calc_flow/runtime.py +979 -0
- calc_flow/store.py +138 -0
- calc_flow/symbolic/__init__.py +65 -0
- calc_flow/symbolic/_generated_rolling_kernels.py +23 -0
- calc_flow/symbolic/analyzer.py +2280 -0
- calc_flow/symbolic/domains.py +77 -0
- calc_flow/symbolic/errors.py +58 -0
- calc_flow/symbolic/expr.py +662 -0
- calc_flow/symbolic/lower/__init__.py +31 -0
- calc_flow/symbolic/lower/planners.py +1270 -0
- calc_flow/symbolic/lower/program.py +840 -0
- calc_flow/symbolic/lower/segments.py +836 -0
- calc_flow/symbolic/lower/strategies.py +1472 -0
- calc_flow/symbolic/nodes.py +603 -0
- calc_flow/symbolic/ops.py +1155 -0
- calc_flow/symbolic/optimizer.py +600 -0
- calc_flow/symbolic/program.py +377 -0
- calc_flow/symbolic/types.py +110 -0
- calc_flow/symbolic/windows.py +153 -0
- calc_flow/udf.py +19 -0
- calc_flow_python-4.0.0.dist-info/METADATA +376 -0
- calc_flow_python-4.0.0.dist-info/RECORD +35 -0
- calc_flow_python-4.0.0.dist-info/WHEEL +4 -0
- calc_flow_python-4.0.0.dist-info/licenses/LICENSE +202 -0
- calc_flow_python-4.0.0.dist-info/sboms/calc-flow-python.cyclonedx.json +10081 -0
|
@@ -0,0 +1,600 @@
|
|
|
1
|
+
"""Common subexpression extraction over resolved row-local forests.
|
|
2
|
+
|
|
3
|
+
Structurally identical, non-trivial subtrees referenced at least twice are
|
|
4
|
+
materialized once as ``__cf_cse_N`` columns in tiered expression nodes so the
|
|
5
|
+
final fused node computes every shared subexpression exactly once. Tiers are
|
|
6
|
+
emitted deepest first; discovery order, naming, and passthrough order are all
|
|
7
|
+
deterministic functions of the declaration.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
from collections.abc import Callable
|
|
13
|
+
from dataclasses import dataclass
|
|
14
|
+
|
|
15
|
+
from calc_flow.symbolic._generated_rolling_kernels import (
|
|
16
|
+
ROLLING_KERNEL_CAPABILITIES,
|
|
17
|
+
)
|
|
18
|
+
from calc_flow.symbolic.nodes import CStr, Node, build
|
|
19
|
+
|
|
20
|
+
_TRIVIAL = frozenset({"column_ref", "literal"})
|
|
21
|
+
_PREDICATE_KEY = ("predicate",)
|
|
22
|
+
_FIXED_TYPE_BYTES = {
|
|
23
|
+
"bool": 1,
|
|
24
|
+
"int8": 1,
|
|
25
|
+
"uint8": 1,
|
|
26
|
+
"int16": 2,
|
|
27
|
+
"uint16": 2,
|
|
28
|
+
"float32": 4,
|
|
29
|
+
"int32": 4,
|
|
30
|
+
"uint32": 4,
|
|
31
|
+
"float64": 8,
|
|
32
|
+
"int64": 8,
|
|
33
|
+
"uint64": 8,
|
|
34
|
+
"timestamp[us]": 8,
|
|
35
|
+
"timestamp[us, UTC]": 8,
|
|
36
|
+
}
|
|
37
|
+
_PRIMITIVE_NUMERIC_TYPES = frozenset(
|
|
38
|
+
{
|
|
39
|
+
"int8",
|
|
40
|
+
"int16",
|
|
41
|
+
"int32",
|
|
42
|
+
"int64",
|
|
43
|
+
"uint8",
|
|
44
|
+
"uint16",
|
|
45
|
+
"uint32",
|
|
46
|
+
"uint64",
|
|
47
|
+
"float32",
|
|
48
|
+
"float64",
|
|
49
|
+
}
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
@dataclass(frozen=True, slots=True)
|
|
54
|
+
class FusedSegment:
|
|
55
|
+
"""The extraction outcome: emission-ordered tiers, final selects, predicate."""
|
|
56
|
+
|
|
57
|
+
tiers: tuple[tuple[tuple[str, Node], ...], ...]
|
|
58
|
+
selects: tuple[tuple[str, Node], ...]
|
|
59
|
+
predicate: Node | None
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def expression_refs(tree: Node, /) -> frozenset[str]:
|
|
63
|
+
"""Collect every column reference name inside one resolved tree."""
|
|
64
|
+
|
|
65
|
+
names: set[str] = set()
|
|
66
|
+
|
|
67
|
+
def walk(node: Node) -> None:
|
|
68
|
+
if node.op.name == "column_ref":
|
|
69
|
+
value = node.attr("name")
|
|
70
|
+
if isinstance(value, CStr):
|
|
71
|
+
names.add(value.value)
|
|
72
|
+
return
|
|
73
|
+
for argument in node.args:
|
|
74
|
+
walk(argument)
|
|
75
|
+
|
|
76
|
+
walk(tree)
|
|
77
|
+
return frozenset(names)
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def _subtree_counts(forest: list[tuple[tuple[str, ...], Node]], /) -> dict[str, int]:
|
|
81
|
+
counts: dict[str, int] = {}
|
|
82
|
+
|
|
83
|
+
def walk(node: Node) -> None:
|
|
84
|
+
counts[node.digest] = counts.get(node.digest, 0) + 1
|
|
85
|
+
for argument in node.args:
|
|
86
|
+
walk(argument)
|
|
87
|
+
|
|
88
|
+
for _, tree in forest:
|
|
89
|
+
walk(tree)
|
|
90
|
+
return counts
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _maximal_candidates(
|
|
94
|
+
forest: list[tuple[tuple[str, ...], Node]],
|
|
95
|
+
counts: dict[str, int],
|
|
96
|
+
/,
|
|
97
|
+
) -> list[tuple[str, Node]]:
|
|
98
|
+
"""Shared non-trivial subtrees not contained in any shared subtree.
|
|
99
|
+
|
|
100
|
+
A subtree that only occurs inside other shared subtrees is deferred: once
|
|
101
|
+
the enclosing candidates are rewritten to references, the deferred subtree
|
|
102
|
+
is rediscovered in a deeper (earlier-emitted) tier, so no materialized
|
|
103
|
+
column ever aliases another.
|
|
104
|
+
"""
|
|
105
|
+
|
|
106
|
+
contained: set[str] = set()
|
|
107
|
+
|
|
108
|
+
def mark_descendants(node: Node) -> None:
|
|
109
|
+
for argument in node.args:
|
|
110
|
+
contained.add(argument.digest)
|
|
111
|
+
mark_descendants(argument)
|
|
112
|
+
|
|
113
|
+
chosen: list[tuple[str, Node]] = []
|
|
114
|
+
seen: set[str] = set()
|
|
115
|
+
|
|
116
|
+
def walk(node: Node) -> None:
|
|
117
|
+
if counts.get(node.digest, 0) >= 2 and node.op.name not in _TRIVIAL:
|
|
118
|
+
if node.digest not in seen:
|
|
119
|
+
seen.add(node.digest)
|
|
120
|
+
chosen.append((node.digest, node))
|
|
121
|
+
mark_descendants(node)
|
|
122
|
+
return
|
|
123
|
+
for argument in node.args:
|
|
124
|
+
walk(argument)
|
|
125
|
+
|
|
126
|
+
for _, tree in forest:
|
|
127
|
+
walk(tree)
|
|
128
|
+
return [(digest, tree) for digest, tree in chosen if digest not in contained]
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
def _rewrite(tree: Node, replacements: dict[str, str], /) -> Node:
|
|
132
|
+
replacement = replacements.get(tree.digest)
|
|
133
|
+
if replacement is not None:
|
|
134
|
+
return build("column_ref", (), {"name": CStr(replacement)})
|
|
135
|
+
if not tree.args:
|
|
136
|
+
return tree
|
|
137
|
+
return build(
|
|
138
|
+
tree.op.name,
|
|
139
|
+
tuple(_rewrite(argument, replacements) for argument in tree.args),
|
|
140
|
+
dict(tree.attrs.entries),
|
|
141
|
+
version=tree.op.version,
|
|
142
|
+
)
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
def extract_common(
|
|
146
|
+
selects: tuple[tuple[str, Node], ...],
|
|
147
|
+
predicate: Node | None,
|
|
148
|
+
reserved: frozenset[str],
|
|
149
|
+
/,
|
|
150
|
+
) -> FusedSegment:
|
|
151
|
+
"""Extract shared subexpressions into emission-ordered materialization tiers.
|
|
152
|
+
|
|
153
|
+
``reserved`` names (declared fields) are never used for materialized
|
|
154
|
+
columns. Discovery-order tiers are reversed for emission so deeper shared
|
|
155
|
+
subexpressions are computed before the tiers that reference them.
|
|
156
|
+
"""
|
|
157
|
+
|
|
158
|
+
forest: list[tuple[tuple[str, ...], Node]] = [
|
|
159
|
+
(("select", name), tree) for name, tree in selects
|
|
160
|
+
]
|
|
161
|
+
if predicate is not None:
|
|
162
|
+
forest.append((_PREDICATE_KEY, predicate))
|
|
163
|
+
iterations: list[tuple[str, ...]] = []
|
|
164
|
+
counter = 0
|
|
165
|
+
|
|
166
|
+
def next_name() -> str:
|
|
167
|
+
nonlocal counter
|
|
168
|
+
while f"__cf_cse_{counter}" in reserved:
|
|
169
|
+
counter += 1
|
|
170
|
+
name = f"__cf_cse_{counter}"
|
|
171
|
+
counter += 1
|
|
172
|
+
return name
|
|
173
|
+
|
|
174
|
+
while True:
|
|
175
|
+
counts = _subtree_counts(forest)
|
|
176
|
+
candidates = _maximal_candidates(forest, counts)
|
|
177
|
+
if not candidates:
|
|
178
|
+
break
|
|
179
|
+
replacements: dict[str, str] = {}
|
|
180
|
+
names: list[str] = []
|
|
181
|
+
defs: list[tuple[tuple[str, ...], Node]] = []
|
|
182
|
+
for digest, tree in candidates:
|
|
183
|
+
name = next_name()
|
|
184
|
+
replacements[digest] = name
|
|
185
|
+
names.append(name)
|
|
186
|
+
defs.append((("cse", name), tree))
|
|
187
|
+
iterations.append(tuple(names))
|
|
188
|
+
forest = [(key, _rewrite(tree, replacements)) for key, tree in forest]
|
|
189
|
+
forest.extend(defs)
|
|
190
|
+
|
|
191
|
+
by_key = {key: tree for key, tree in forest}
|
|
192
|
+
tiers = tuple(
|
|
193
|
+
tuple((name, by_key[("cse", name)]) for name in names)
|
|
194
|
+
for names in reversed(iterations)
|
|
195
|
+
)
|
|
196
|
+
final_selects = tuple((name, by_key[("select", name)]) for name, _ in selects)
|
|
197
|
+
final_predicate = by_key.get(_PREDICATE_KEY)
|
|
198
|
+
return FusedSegment(tiers, final_selects, final_predicate)
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
def _document_nodes(document: dict[str, object], /) -> list[dict[str, object]]:
|
|
202
|
+
graph = document.get("graph")
|
|
203
|
+
if not isinstance(graph, dict):
|
|
204
|
+
return []
|
|
205
|
+
nodes = graph.get("nodes")
|
|
206
|
+
return nodes if isinstance(nodes, list) else [] # type: ignore[return-value]
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def _input_schema_fields(
|
|
210
|
+
node: dict[str, object], index: int = 0, /
|
|
211
|
+
) -> list[dict[str, object]]:
|
|
212
|
+
input_ports = node.get("input_ports")
|
|
213
|
+
if not isinstance(input_ports, list):
|
|
214
|
+
return []
|
|
215
|
+
if not input_ports:
|
|
216
|
+
return []
|
|
217
|
+
if index >= len(input_ports):
|
|
218
|
+
return []
|
|
219
|
+
port = input_ports[index]
|
|
220
|
+
if not isinstance(port, dict):
|
|
221
|
+
return []
|
|
222
|
+
schema = port.get("schema")
|
|
223
|
+
if not isinstance(schema, list):
|
|
224
|
+
return []
|
|
225
|
+
return [field for field in schema if isinstance(field, dict)]
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
def _state_layout(node: dict[str, object], index: int = 0, /) -> str:
|
|
229
|
+
fields = _input_schema_fields(node, index)
|
|
230
|
+
fixed_bytes = sum(
|
|
231
|
+
_FIXED_TYPE_BYTES.get(field.get("data_type"), 0) for field in fields
|
|
232
|
+
)
|
|
233
|
+
variable_columns = sum(
|
|
234
|
+
field.get("data_type") not in _FIXED_TYPE_BYTES for field in fields
|
|
235
|
+
)
|
|
236
|
+
layout = (
|
|
237
|
+
f"retained_columns={len(fields)} fixed_bytes_per_row={fixed_bytes}"
|
|
238
|
+
f" variable_columns={variable_columns}"
|
|
239
|
+
)
|
|
240
|
+
return layout
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
def _frame_bounds(frame: object, /) -> tuple[tuple[int, ...], tuple[int, ...]]:
|
|
244
|
+
if not isinstance(frame, dict):
|
|
245
|
+
return (), ()
|
|
246
|
+
kind = frame.get("kind")
|
|
247
|
+
if kind == "rows":
|
|
248
|
+
size = frame.get("size")
|
|
249
|
+
return ((size,) if isinstance(size, int) else ()), ()
|
|
250
|
+
if kind == "duration":
|
|
251
|
+
micros = frame.get("micros")
|
|
252
|
+
return (), ((micros,) if isinstance(micros, int) else ())
|
|
253
|
+
return (), ()
|
|
254
|
+
|
|
255
|
+
|
|
256
|
+
def _rolling_output_bounds(
|
|
257
|
+
output: object, /
|
|
258
|
+
) -> tuple[tuple[int, ...], tuple[int, ...]]:
|
|
259
|
+
if not isinstance(output, dict):
|
|
260
|
+
return (), ()
|
|
261
|
+
if output.get("kind") == "difference":
|
|
262
|
+
left_rows, left_durations = _rolling_output_bounds(output.get("left"))
|
|
263
|
+
right_rows, right_durations = _rolling_output_bounds(output.get("right"))
|
|
264
|
+
return (*left_rows, *right_rows), (*left_durations, *right_durations)
|
|
265
|
+
periods = output.get("periods")
|
|
266
|
+
row_bounds = (periods + 1,) if isinstance(periods, int) else ()
|
|
267
|
+
frame_rows, duration_bounds = _frame_bounds(output.get("frame"))
|
|
268
|
+
return (*row_bounds, *frame_rows), duration_bounds
|
|
269
|
+
|
|
270
|
+
|
|
271
|
+
def _rolling_boundary(outputs: list[object], /) -> str:
|
|
272
|
+
row_bounds: list[int] = []
|
|
273
|
+
duration_bounds: list[int] = []
|
|
274
|
+
for output in outputs:
|
|
275
|
+
output_rows, output_durations = _rolling_output_bounds(output)
|
|
276
|
+
row_bounds.extend(output_rows)
|
|
277
|
+
duration_bounds.extend(output_durations)
|
|
278
|
+
bounds = []
|
|
279
|
+
if row_bounds:
|
|
280
|
+
bounds.append(f"rows={max(row_bounds)}")
|
|
281
|
+
if duration_bounds:
|
|
282
|
+
bounds.append(f"duration_micros={max(duration_bounds)}")
|
|
283
|
+
return " ".join(bounds) or "constant"
|
|
284
|
+
|
|
285
|
+
|
|
286
|
+
def _cross_section_boundary(spec: dict[str, object], /) -> str:
|
|
287
|
+
grouping = spec.get("grouping")
|
|
288
|
+
if isinstance(grouping, dict) and grouping.get("kind") == "fixed_bucket":
|
|
289
|
+
return f"bucket_width_micros={grouping.get('width_micros')}"
|
|
290
|
+
return "exact_time_groups"
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
def _state_cost(node: dict[str, object], /) -> str | None:
|
|
294
|
+
operator = node.get("operator")
|
|
295
|
+
spec = operator.get("spec") if isinstance(operator, dict) else None
|
|
296
|
+
kind = operator.get("kind") if isinstance(operator, dict) else None
|
|
297
|
+
if kind == "stream_join" and isinstance(spec, dict):
|
|
298
|
+
limits = spec.get("limits")
|
|
299
|
+
if not isinstance(limits, dict):
|
|
300
|
+
return None
|
|
301
|
+
return (
|
|
302
|
+
f" state {node['id']}"
|
|
303
|
+
f" max_state_rows_per_side={limits.get('max_state_rows_per_side')}"
|
|
304
|
+
f" max_state_bytes_per_side={limits.get('max_state_bytes_per_side')}"
|
|
305
|
+
" max_matches_per_input_batch="
|
|
306
|
+
f"{limits.get('max_matches_per_input_batch')}"
|
|
307
|
+
f" left_{_state_layout(node, 0)} right_{_state_layout(node, 1)}"
|
|
308
|
+
)
|
|
309
|
+
outputs = spec.get("outputs") if isinstance(spec, dict) else None
|
|
310
|
+
if not isinstance(outputs, list):
|
|
311
|
+
return None
|
|
312
|
+
layout = _state_layout(node)
|
|
313
|
+
if kind == "rolling":
|
|
314
|
+
return f" state {node['id']} {_rolling_boundary(outputs)} {layout}"
|
|
315
|
+
if kind == "cross_section":
|
|
316
|
+
boundary = _cross_section_boundary(spec)
|
|
317
|
+
return f" state {node['id']} {boundary} active_groups=runtime {layout}"
|
|
318
|
+
return None
|
|
319
|
+
|
|
320
|
+
|
|
321
|
+
def _static_array_weight(declaration: object, /) -> tuple[bool, int | None]:
|
|
322
|
+
if not isinstance(declaration, dict):
|
|
323
|
+
return False, None
|
|
324
|
+
if declaration.get("kind") != "array":
|
|
325
|
+
return False, None
|
|
326
|
+
shape = declaration.get("shape")
|
|
327
|
+
width = _FIXED_TYPE_BYTES.get(declaration.get("dtype"))
|
|
328
|
+
if not isinstance(shape, list):
|
|
329
|
+
return False, None
|
|
330
|
+
if width is None:
|
|
331
|
+
return False, None
|
|
332
|
+
elements = 1
|
|
333
|
+
for dimension in shape:
|
|
334
|
+
if not isinstance(dimension, int):
|
|
335
|
+
return True, None
|
|
336
|
+
elements *= dimension
|
|
337
|
+
return True, elements * width
|
|
338
|
+
|
|
339
|
+
|
|
340
|
+
def _static_weight_bytes(document: dict[str, object], /) -> int | None:
|
|
341
|
+
static_inputs = document.get("static_inputs")
|
|
342
|
+
if not isinstance(static_inputs, list):
|
|
343
|
+
return None
|
|
344
|
+
for declaration in static_inputs:
|
|
345
|
+
found, weight = _static_array_weight(declaration)
|
|
346
|
+
if found:
|
|
347
|
+
return weight
|
|
348
|
+
return None
|
|
349
|
+
|
|
350
|
+
|
|
351
|
+
def _copy_cost(
|
|
352
|
+
node: dict[str, object], static_weight_bytes: int | None, /
|
|
353
|
+
) -> str | None:
|
|
354
|
+
operator = node.get("operator")
|
|
355
|
+
if not isinstance(operator, dict) or operator.get("kind") != "external":
|
|
356
|
+
return None
|
|
357
|
+
options = operator.get("options")
|
|
358
|
+
columns = options.get("columns") if isinstance(options, dict) else None
|
|
359
|
+
column_count = len(columns) if isinstance(columns, list) else 0
|
|
360
|
+
backend = operator.get("provider")
|
|
361
|
+
device_copy = "yes" if backend == "jax" else "no"
|
|
362
|
+
weights = "runtime" if static_weight_bytes is None else str(static_weight_bytes)
|
|
363
|
+
return (
|
|
364
|
+
f" copies {node['id']} table_to_dense columns={column_count}"
|
|
365
|
+
f" rows=runtime host_to_device={device_copy} static_weights_bytes={weights}"
|
|
366
|
+
)
|
|
367
|
+
|
|
368
|
+
|
|
369
|
+
def _provider_cost(node: dict[str, object], /) -> str | None:
|
|
370
|
+
operator = node.get("operator")
|
|
371
|
+
if not isinstance(operator, dict) or operator.get("kind") != "external":
|
|
372
|
+
return None
|
|
373
|
+
return (
|
|
374
|
+
f" providers {node['id']} {operator.get('provider')}:"
|
|
375
|
+
f"{operator.get('name')}@{operator.get('version')} calls_per_microbatch=1"
|
|
376
|
+
)
|
|
377
|
+
|
|
378
|
+
|
|
379
|
+
def _nodes_of_kind(
|
|
380
|
+
nodes: list[dict[str, object]], kind: str, /
|
|
381
|
+
) -> list[dict[str, object]]:
|
|
382
|
+
return [
|
|
383
|
+
node
|
|
384
|
+
for node in nodes
|
|
385
|
+
if isinstance(node.get("operator"), dict)
|
|
386
|
+
and node["operator"].get("kind") == kind # type: ignore[union-attr]
|
|
387
|
+
]
|
|
388
|
+
|
|
389
|
+
|
|
390
|
+
def _state_output_count(items: list[dict[str, object]], /) -> int:
|
|
391
|
+
count = 0
|
|
392
|
+
for item in items:
|
|
393
|
+
operator = item["operator"]
|
|
394
|
+
spec = operator.get("spec") if isinstance(operator, dict) else None
|
|
395
|
+
outputs = spec.get("outputs") if isinstance(spec, dict) else None
|
|
396
|
+
count += len(outputs) if isinstance(outputs, list) else 0
|
|
397
|
+
return count
|
|
398
|
+
|
|
399
|
+
|
|
400
|
+
def _rolling_fusion_count(items: list[dict[str, object]], /) -> int:
|
|
401
|
+
count = 0
|
|
402
|
+
for item in items:
|
|
403
|
+
operator = item["operator"]
|
|
404
|
+
spec = operator.get("spec") if isinstance(operator, dict) else None
|
|
405
|
+
outputs = spec.get("outputs") if isinstance(spec, dict) else None
|
|
406
|
+
if isinstance(outputs, list):
|
|
407
|
+
count += sum(
|
|
408
|
+
isinstance(output, dict) and output.get("kind") == "difference"
|
|
409
|
+
for output in outputs
|
|
410
|
+
)
|
|
411
|
+
return count
|
|
412
|
+
|
|
413
|
+
|
|
414
|
+
def _rolling_leaf_outputs(output: object, /) -> tuple[dict[str, object], ...]:
|
|
415
|
+
if not isinstance(output, dict):
|
|
416
|
+
return ()
|
|
417
|
+
if output.get("kind") != "difference":
|
|
418
|
+
return (output,)
|
|
419
|
+
return (
|
|
420
|
+
*_rolling_leaf_outputs(output.get("left")),
|
|
421
|
+
*_rolling_leaf_outputs(output.get("right")),
|
|
422
|
+
)
|
|
423
|
+
|
|
424
|
+
|
|
425
|
+
def _rolling_frame_key(output: dict[str, object], /) -> tuple[object, object]:
|
|
426
|
+
frame = output.get("frame")
|
|
427
|
+
if not isinstance(frame, dict):
|
|
428
|
+
return (None, None)
|
|
429
|
+
kind = frame.get("kind")
|
|
430
|
+
coordinate = frame.get("size") if kind == "rows" else frame.get("micros")
|
|
431
|
+
return kind, coordinate
|
|
432
|
+
|
|
433
|
+
|
|
434
|
+
def _rolling_group_key(output: dict[str, object], /) -> tuple[object, ...] | None:
|
|
435
|
+
kind = output.get("kind")
|
|
436
|
+
frame = _rolling_frame_key(output)
|
|
437
|
+
if kind in {"count", "sum", "mean", "variance", "stddev"}:
|
|
438
|
+
return "numeric", output.get("input"), *frame
|
|
439
|
+
if kind in {"min", "max"}:
|
|
440
|
+
return "extrema", kind, output.get("input"), *frame
|
|
441
|
+
if kind in {"covariance", "correlation"}:
|
|
442
|
+
return "pair", output.get("left"), output.get("right"), *frame
|
|
443
|
+
if kind == "ewma":
|
|
444
|
+
return "ewma", output.get("input"), output.get("span")
|
|
445
|
+
return None
|
|
446
|
+
|
|
447
|
+
|
|
448
|
+
def _first_rolling_fallback(
|
|
449
|
+
outputs: tuple[dict[str, object], ...], field_types: dict[str, object], /
|
|
450
|
+
) -> str | None:
|
|
451
|
+
for output in outputs:
|
|
452
|
+
fallback = _rolling_kernel_fallback(output, field_types)
|
|
453
|
+
if fallback is not None:
|
|
454
|
+
return fallback
|
|
455
|
+
return None
|
|
456
|
+
|
|
457
|
+
|
|
458
|
+
def _rolling_input_columns(
|
|
459
|
+
output: dict[str, object], transition: object, /
|
|
460
|
+
) -> tuple[object, ...]:
|
|
461
|
+
if transition == "pair":
|
|
462
|
+
return output.get("left"), output.get("right")
|
|
463
|
+
return (output.get("input"),)
|
|
464
|
+
|
|
465
|
+
|
|
466
|
+
def _rolling_numeric_fallback(
|
|
467
|
+
kind: object,
|
|
468
|
+
columns: tuple[object, ...],
|
|
469
|
+
field_types: dict[str, object],
|
|
470
|
+
/,
|
|
471
|
+
) -> str | None:
|
|
472
|
+
for column in columns:
|
|
473
|
+
data_type = field_types.get(column) if isinstance(column, str) else None
|
|
474
|
+
if data_type not in _PRIMITIVE_NUMERIC_TYPES:
|
|
475
|
+
return f"primitive_{kind}_requires_numeric_column_{column}"
|
|
476
|
+
return None
|
|
477
|
+
|
|
478
|
+
|
|
479
|
+
def _rolling_kernel_fallback(
|
|
480
|
+
output: dict[str, object], field_types: dict[str, object], /
|
|
481
|
+
) -> str | None:
|
|
482
|
+
kind = output.get("kind")
|
|
483
|
+
capability = ROLLING_KERNEL_CAPABILITIES.get(kind)
|
|
484
|
+
if capability is None:
|
|
485
|
+
return f"primitive_{kind}_missing_from_census"
|
|
486
|
+
transition = capability[0]
|
|
487
|
+
if transition is None:
|
|
488
|
+
return f"primitive_{kind}_has_no_typed_transition"
|
|
489
|
+
if kind == "difference":
|
|
490
|
+
return _first_rolling_fallback(_rolling_leaf_outputs(output), field_types)
|
|
491
|
+
return _rolling_numeric_fallback(
|
|
492
|
+
kind, _rolling_input_columns(output, transition), field_types
|
|
493
|
+
)
|
|
494
|
+
|
|
495
|
+
|
|
496
|
+
def _rolling_spec_outputs(
|
|
497
|
+
node: dict[str, object], /
|
|
498
|
+
) -> tuple[dict[str, object], tuple[dict[str, object], ...]] | None:
|
|
499
|
+
operator = node.get("operator")
|
|
500
|
+
spec = operator.get("spec") if isinstance(operator, dict) else None
|
|
501
|
+
raw_outputs = spec.get("outputs") if isinstance(spec, dict) else None
|
|
502
|
+
if not isinstance(spec, dict) or not isinstance(raw_outputs, list):
|
|
503
|
+
return None
|
|
504
|
+
return spec, tuple(output for output in raw_outputs if isinstance(output, dict))
|
|
505
|
+
|
|
506
|
+
|
|
507
|
+
def _rolling_field_types(node: dict[str, object], /) -> dict[str, object]:
|
|
508
|
+
return {
|
|
509
|
+
str(field.get("name")): field.get("data_type")
|
|
510
|
+
for field in _input_schema_fields(node)
|
|
511
|
+
}
|
|
512
|
+
|
|
513
|
+
|
|
514
|
+
def _rolling_state_groups(
|
|
515
|
+
outputs: tuple[dict[str, object], ...], /
|
|
516
|
+
) -> set[tuple[object, ...]]:
|
|
517
|
+
groups: set[tuple[object, ...]] = set()
|
|
518
|
+
for output in outputs:
|
|
519
|
+
for leaf in _rolling_leaf_outputs(output):
|
|
520
|
+
key = _rolling_group_key(leaf)
|
|
521
|
+
if key is not None:
|
|
522
|
+
groups.add(key)
|
|
523
|
+
return groups
|
|
524
|
+
|
|
525
|
+
|
|
526
|
+
def _rolling_order(spec: dict[str, object], /) -> str:
|
|
527
|
+
values = (
|
|
528
|
+
spec.get("event_time"),
|
|
529
|
+
*(spec.get("partition_by") or []),
|
|
530
|
+
*(spec.get("sequence_by") or []),
|
|
531
|
+
)
|
|
532
|
+
return ",".join(str(value) for value in values)
|
|
533
|
+
|
|
534
|
+
|
|
535
|
+
def _rolling_kernel_line(node: dict[str, object], /) -> str | None:
|
|
536
|
+
plan = _rolling_spec_outputs(node)
|
|
537
|
+
if plan is None:
|
|
538
|
+
return None
|
|
539
|
+
spec, outputs = plan
|
|
540
|
+
groups = _rolling_state_groups(outputs)
|
|
541
|
+
fallback = _first_rolling_fallback(outputs, _rolling_field_types(node))
|
|
542
|
+
selected = "ordered_primitive" if fallback is None and groups else "general"
|
|
543
|
+
complexity = "amortized_constant" if selected == "ordered_primitive" else "general"
|
|
544
|
+
profile = spec.get("numerical_profile", "stable_v1")
|
|
545
|
+
return (
|
|
546
|
+
f" rolling kernel {node['id']} selected={selected}"
|
|
547
|
+
f" profile={profile} complexity={complexity} order={_rolling_order(spec)}"
|
|
548
|
+
f" shared_state_groups={len(groups)} fallback={fallback or 'none'}"
|
|
549
|
+
)
|
|
550
|
+
|
|
551
|
+
|
|
552
|
+
def _cost_lines(
|
|
553
|
+
nodes: list[dict[str, object]],
|
|
554
|
+
renderer: Callable[[dict[str, object]], str | None],
|
|
555
|
+
/,
|
|
556
|
+
) -> tuple[str, ...]:
|
|
557
|
+
lines: list[str] = []
|
|
558
|
+
for node in nodes:
|
|
559
|
+
line = renderer(node)
|
|
560
|
+
if line is not None:
|
|
561
|
+
lines.append(line)
|
|
562
|
+
return tuple(lines)
|
|
563
|
+
|
|
564
|
+
|
|
565
|
+
def explain_optimization(document: dict[str, object], /) -> tuple[str, ...]:
|
|
566
|
+
"""Render deterministic physical sharing and bounded cost facts."""
|
|
567
|
+
|
|
568
|
+
nodes = _document_nodes(document)
|
|
569
|
+
cse_count = sum("__cf_cse_" in str(node.get("id")) for node in nodes)
|
|
570
|
+
rolling = _nodes_of_kind(nodes, "rolling")
|
|
571
|
+
cross_section = _nodes_of_kind(nodes, "cross_section")
|
|
572
|
+
stream_join = _nodes_of_kind(nodes, "stream_join")
|
|
573
|
+
external = _nodes_of_kind(nodes, "external")
|
|
574
|
+
|
|
575
|
+
lines = (
|
|
576
|
+
" optimization",
|
|
577
|
+
f" cse materializations {cse_count}",
|
|
578
|
+
" rolling state_stages"
|
|
579
|
+
f" {len(rolling)} shared_outputs {_state_output_count(rolling)}",
|
|
580
|
+
" rolling fused_outputs"
|
|
581
|
+
f" {_rolling_fusion_count(rolling)} hidden_materializations 0",
|
|
582
|
+
" cross_section grouping_stages"
|
|
583
|
+
f" {len(cross_section)} shared_outputs {_state_output_count(cross_section)}",
|
|
584
|
+
f" stream_join state_stages {len(stream_join)}",
|
|
585
|
+
f" array fused_stages {len(external)} provider_calls_per_microbatch"
|
|
586
|
+
f" {len(external)}",
|
|
587
|
+
)
|
|
588
|
+
kernels = _cost_lines(rolling, _rolling_kernel_line)
|
|
589
|
+
state = _cost_lines(nodes, _state_cost)
|
|
590
|
+
static_weight_bytes = _static_weight_bytes(document)
|
|
591
|
+
copies = _cost_lines(nodes, lambda node: _copy_cost(node, static_weight_bytes))
|
|
592
|
+
providers = _cost_lines(nodes, _provider_cost)
|
|
593
|
+
return (
|
|
594
|
+
*lines,
|
|
595
|
+
*(kernels or (" rolling kernels none",)),
|
|
596
|
+
" costs",
|
|
597
|
+
*(state or (" state none",)),
|
|
598
|
+
*(copies or (" copies none",)),
|
|
599
|
+
*(providers or (" providers none",)),
|
|
600
|
+
)
|