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.
@@ -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
+ )