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,1270 @@
1
+ """Rolling and cross-section stage planning for the symbolic lowerer.
2
+
3
+ Moved verbatim from ``symbolic/lower.py``."""
4
+
5
+ from __future__ import annotations
6
+
7
+ from dataclasses import dataclass, replace
8
+ from typing import TYPE_CHECKING
9
+
10
+ from calc_flow.pipeline import (
11
+ _canonical,
12
+ )
13
+ from calc_flow.symbolic import errors
14
+ from calc_flow.symbolic.analyzer import (
15
+ _rolling_output_type,
16
+ _schema_fields,
17
+ )
18
+ from calc_flow.symbolic.lower.segments import (
19
+ _CROSS_SECTION_DDOF,
20
+ _CROSS_SECTION_ORDERING,
21
+ _CROSS_SECTION_PRIMITIVES,
22
+ _ROLLING_DDOF_PRIMITIVES,
23
+ _ROLLING_PAIR_PRIMITIVES,
24
+ _ROLLING_PRIMITIVES,
25
+ _cbool,
26
+ _cint,
27
+ _cnumber,
28
+ _cstr,
29
+ _cstr_seq,
30
+ _expression_node,
31
+ _field_json,
32
+ _find_cross_section,
33
+ _find_ready_rolling,
34
+ _find_rolling,
35
+ _fused_difference_outputs,
36
+ _fused_float_leaf,
37
+ _plan_stateful_inputs,
38
+ _quote_identifier,
39
+ _replace_materialized,
40
+ _rolling_declaration_requires_ewma,
41
+ _rolling_frame,
42
+ _RollingPipeline,
43
+ _RollingPlan,
44
+ _Segment,
45
+ _StatefulInputRequest,
46
+ )
47
+ from calc_flow.symbolic.nodes import (
48
+ CEnum,
49
+ CMap,
50
+ Node,
51
+ )
52
+ from calc_flow.symbolic.types import Field
53
+
54
+ if TYPE_CHECKING:
55
+ from calc_flow.symbolic.program import Program
56
+
57
+
58
+ # Rolling planning validates every lag/delta/aggregate occurrence with
59
+ # stable, declaration-ordered error paths before emitting the frozen node
60
+ # shape.
61
+ def _plan_rolling_stage(
62
+ output_name: str,
63
+ segment: _Segment,
64
+ path: str,
65
+ allowed_lateness_micros: int,
66
+ late_policy: str,
67
+ occurrences: tuple[Node, ...],
68
+ input_fields: tuple[Field, ...],
69
+ stage_number: int,
70
+ stage_count: int,
71
+ /,
72
+ ) -> _RollingPlan:
73
+ # #lizard forgives
74
+ input_types = {field.name: field for field in input_fields}
75
+ entity_by = _cstr_seq(segment.input_node.attr("entity_by"))
76
+ sequence_by = _cstr_seq(segment.input_node.attr("sequence_by"))
77
+ event_time = _cstr(segment.input_node.attr("event_time"))
78
+ if not entity_by or not sequence_by or not event_time:
79
+ errors.raise_compile(
80
+ path,
81
+ errors.ORDERING_REQUIRED,
82
+ "rolling temporal primitives require declared entity_by,"
83
+ " event_time, and sequence_by ordering keys on the input table",
84
+ )
85
+
86
+ whole_feature = {
87
+ tree.digest: name
88
+ for name, tree in segment.env
89
+ if tree.op.name in _ROLLING_PRIMITIVES
90
+ }
91
+ fused_differences = _fused_difference_outputs(segment, occurrences)
92
+ fused_leaf_digests = {
93
+ argument.digest for _, tree in fused_differences for argument in tree.args
94
+ }
95
+ fused_roots = {tree.digest: name for name, tree in fused_differences}
96
+ unfused_leaf_digests = {
97
+ subtree.digest
98
+ for _, tree in segment.env
99
+ for subtree in _find_rolling(_replace_materialized(tree, fused_roots))
100
+ }
101
+ hidden_fused_leaf_digests = fused_leaf_digests - unfused_leaf_digests
102
+ used_names = set(input_types) | {name for name, _ in segment.env}
103
+ stage_fragment = "" if stage_count == 1 else f"{stage_number}_"
104
+ stage_suffix = "" if stage_count == 1 else f"_{stage_number}"
105
+ materialization = _plan_stateful_inputs(
106
+ _StatefulInputRequest(
107
+ output_name,
108
+ path,
109
+ f"roll_input_{stage_fragment}".removesuffix("_"),
110
+ f"{output_name}__cf_rolling_input{stage_suffix}",
111
+ "rolling",
112
+ ),
113
+ input_fields,
114
+ used_names,
115
+ tuple(
116
+ (subtree.op.name, argument)
117
+ for subtree in occurrences
118
+ for argument in subtree.args
119
+ ),
120
+ )
121
+ materializations = dict(materialization.names)
122
+ state_input_fields = materialization.input_fields
123
+ input_types = {field.name: field for field in state_input_fields}
124
+ used_names = set(materialization.used_names)
125
+ replacements: dict[str, str] = {}
126
+ declarations: list[dict[str, object]] = []
127
+ derived_fields: list[Field] = []
128
+ for index, subtree in enumerate(occurrences):
129
+ if subtree.digest in hidden_fused_leaf_digests:
130
+ continue
131
+ kind = subtree.op.name
132
+ name = whole_feature.get(subtree.digest)
133
+ if name is None:
134
+ name = f"{output_name}__cf_roll_{stage_fragment}{index}"
135
+ if name in used_names:
136
+ errors.raise_compile(
137
+ f"{path}.{name}",
138
+ errors.DUPLICATE_NAME,
139
+ f"materialized rolling column {name!r} collides with a"
140
+ " declared field",
141
+ )
142
+ used_names.add(name)
143
+ replacements[subtree.digest] = name
144
+ operands: list[tuple[str, str, Field]] = []
145
+ for role, argument in zip(
146
+ ("input", "left", "right"), subtree.args, strict=False
147
+ ):
148
+ input_name = (
149
+ _cstr(argument.attr("name"))
150
+ if argument.op.name == "column_ref"
151
+ else materializations[argument.digest]
152
+ )
153
+ field = input_types.get(input_name)
154
+ if field is None:
155
+ errors.raise_compile(
156
+ f"{path}.{name}",
157
+ errors.SCHEMA_MISMATCH,
158
+ f"rolling {kind} argument column {input_name!r} is not in the"
159
+ " input schema",
160
+ )
161
+ operands.append((role, input_name, field))
162
+ periods = _cint(subtree.attr("periods"))
163
+ if periods is not None:
164
+ declarations.append(
165
+ {
166
+ "kind": kind,
167
+ "primitive_version": 1,
168
+ "input": operands[0][1],
169
+ "output": name,
170
+ "periods": periods,
171
+ }
172
+ )
173
+ derived_fields.append(Field(name, operands[0][2].data_type, nullable=True))
174
+ continue
175
+ if kind == "ewma":
176
+ declarations.append(
177
+ {
178
+ "kind": kind,
179
+ "primitive_version": 1,
180
+ "input": operands[0][1],
181
+ "output": name,
182
+ "span": _cint(subtree.attr("span")),
183
+ "min_periods": _cint(subtree.attr("min_periods")) or 1,
184
+ }
185
+ )
186
+ derived_fields.append(Field(name, "float64", nullable=True))
187
+ continue
188
+ frame = _rolling_frame(subtree, f"{path}.{name}", kind)
189
+ declaration: dict[str, object] = {
190
+ "kind": kind,
191
+ "primitive_version": 1,
192
+ "output": name,
193
+ "frame": frame,
194
+ "min_periods": _cint(subtree.attr("min_periods")) or 1,
195
+ }
196
+ if kind in _ROLLING_PAIR_PRIMITIVES:
197
+ declaration["left"] = operands[0][1]
198
+ declaration["right"] = operands[1][1]
199
+ else:
200
+ declaration["input"] = operands[0][1]
201
+ if kind in _ROLLING_DDOF_PRIMITIVES:
202
+ ddof = _cint(subtree.attr("ddof"))
203
+ declaration["ddof"] = 1 if ddof is None else ddof
204
+ declarations.append(declaration)
205
+ derived_fields.append(
206
+ Field(
207
+ name,
208
+ _rolling_output_type(kind, operands[0][2].data_type) or "float64",
209
+ nullable=True,
210
+ )
211
+ )
212
+
213
+ for name, tree in fused_differences:
214
+ replacements[tree.digest] = name
215
+ declarations.append(
216
+ {
217
+ "kind": "difference",
218
+ "primitive_version": 1,
219
+ "left": _fused_float_leaf(
220
+ tree.args[0],
221
+ materializations,
222
+ input_types,
223
+ f"{path}.{name}.left",
224
+ ),
225
+ "right": _fused_float_leaf(
226
+ tree.args[1],
227
+ materializations,
228
+ input_types,
229
+ f"{path}.{name}.right",
230
+ ),
231
+ "output": name,
232
+ }
233
+ )
234
+ derived_fields.append(Field(name, "float64", nullable=True))
235
+
236
+ node_id = f"{output_name}__cf_rolling{stage_suffix}"
237
+ node: dict[str, object] = {
238
+ "id": node_id,
239
+ "operator": {
240
+ "kind": "rolling",
241
+ "spec": {
242
+ "configuration_version": 1,
243
+ "state_layout_version": (
244
+ 2
245
+ if any(
246
+ _rolling_declaration_requires_ewma(item)
247
+ for item in declarations
248
+ )
249
+ else 1
250
+ ),
251
+ "partition_by": list(entity_by),
252
+ "event_time": event_time,
253
+ "sequence_by": list(sequence_by),
254
+ "outputs": declarations,
255
+ "allowed_lateness_micros": allowed_lateness_micros,
256
+ "late_policy": (
257
+ {"kind": "error", "scope": "envelope"}
258
+ if late_policy == "error"
259
+ else {"kind": "drop", "metrics_version": 1}
260
+ ),
261
+ "value_policy": "stateful_numeric_v1",
262
+ },
263
+ },
264
+ "input_ports": [
265
+ {
266
+ "name": "input",
267
+ "kind": "table",
268
+ "required": True,
269
+ "schema": [_field_json(field) for field in state_input_fields],
270
+ }
271
+ ],
272
+ "output_ports": [
273
+ {
274
+ "name": "output",
275
+ "kind": "table",
276
+ "required": True,
277
+ "schema": [
278
+ *(_field_json(field) for field in state_input_fields),
279
+ *(_field_json(field) for field in derived_fields),
280
+ ],
281
+ }
282
+ ],
283
+ }
284
+ env = tuple(
285
+ (name, _replace_materialized(tree, replacements)) for name, tree in segment.env
286
+ )
287
+ post_predicate = (
288
+ None
289
+ if segment.post_predicate is None
290
+ else _replace_materialized(segment.post_predicate, replacements)
291
+ )
292
+ return _RollingPlan(
293
+ node_id=node_id,
294
+ node=node,
295
+ materialization_node_id=materialization.node_id,
296
+ materialization_node=materialization.node,
297
+ env=env,
298
+ post_predicate=post_predicate,
299
+ input_field_names=(*input_types, *(field.name for field in derived_fields)),
300
+ output_fields=(*state_input_fields, *derived_fields),
301
+ replacements=tuple(replacements.items()),
302
+ )
303
+
304
+
305
+ def _ready_rolling_occurrences(segment: _Segment, /) -> tuple[Node, ...]:
306
+ ordered_trees = [tree for _, tree in segment.env]
307
+ ordered_trees += [
308
+ tree for tree in (segment.predicate, segment.post_predicate) if tree is not None
309
+ ]
310
+ occurrences: list[Node] = []
311
+ seen: set[str] = set()
312
+ for tree in ordered_trees:
313
+ for subtree in _find_ready_rolling(tree):
314
+ if subtree.digest not in seen:
315
+ seen.add(subtree.digest)
316
+ occurrences.append(subtree)
317
+ return tuple(occurrences)
318
+
319
+
320
+ def _rolling_depth(node: Node, /) -> int:
321
+ child_depth = max((_rolling_depth(argument) for argument in node.args), default=0)
322
+ return child_depth + 1 if node.op.name in _ROLLING_PRIMITIVES else child_depth
323
+
324
+
325
+ def _rolling_stage_count(segment: _Segment, /) -> int:
326
+ """Count the deterministic innermost-first rolling layers."""
327
+
328
+ trees = [tree for _, tree in segment.env]
329
+ trees += [
330
+ tree for tree in (segment.predicate, segment.post_predicate) if tree is not None
331
+ ]
332
+ return max((_rolling_depth(tree) for tree in trees), default=0)
333
+
334
+
335
+ def _plan_rolling(
336
+ output_name: str,
337
+ segment: _Segment,
338
+ path: str,
339
+ allowed_lateness_micros: int,
340
+ late_policy: str,
341
+ /,
342
+ ) -> _RollingPipeline | None:
343
+ """Plan every innermost-first rolling layer for one output branch."""
344
+
345
+ stage_count = _rolling_stage_count(segment)
346
+ if stage_count == 0:
347
+ return None
348
+ stages: list[_RollingPlan] = []
349
+ current = segment
350
+ input_fields = _schema_fields(segment.input_node.attr("schema"))
351
+ for stage_number in range(1, stage_count + 1):
352
+ occurrences = _ready_rolling_occurrences(current)
353
+ if not occurrences:
354
+ raise RuntimeError("rolling stage count diverged during lowering")
355
+ stage = _plan_rolling_stage(
356
+ output_name,
357
+ current,
358
+ path,
359
+ allowed_lateness_micros,
360
+ late_policy,
361
+ occurrences,
362
+ input_fields,
363
+ stage_number,
364
+ stage_count,
365
+ )
366
+ stages.append(stage)
367
+ current = replace(
368
+ current,
369
+ env=stage.env,
370
+ post_predicate=stage.post_predicate,
371
+ )
372
+ input_fields = stage.output_fields
373
+ final = stages[-1]
374
+ return _RollingPipeline(
375
+ tuple(stages),
376
+ final.env,
377
+ final.post_predicate,
378
+ final.input_field_names,
379
+ final.output_fields,
380
+ )
381
+
382
+
383
+ @dataclass(frozen=True, slots=True)
384
+ class _CrossSectionPlan:
385
+ """One lowered cross-section stage: the project node plus the rewritten
386
+ row-local environment that references its output columns."""
387
+
388
+ node_id: str
389
+ node: dict[str, object]
390
+ materialization_node_id: str | None
391
+ materialization_node: dict[str, object] | None
392
+ env: tuple[tuple[str, Node], ...]
393
+ post_predicate: Node | None
394
+ input_field_names: tuple[str, ...]
395
+ output_fields: tuple[Field, ...]
396
+ materializations: tuple[tuple[str, str], ...]
397
+ replacements: tuple[tuple[str, str], ...]
398
+
399
+
400
+ def _cross_section_grouping(subtree: Node, path: str, /) -> dict[str, object]:
401
+ """Render the frozen exact-time or fixed-bucket grouping JSON."""
402
+
403
+ grouping = subtree.attr("grouping")
404
+ if isinstance(grouping, CEnum) and grouping.variant == "exact_time":
405
+ return {"kind": "exact_time"}
406
+ if isinstance(grouping, CMap):
407
+ tag = grouping.get("grouping")
408
+ width = _cint(grouping.get("width_micros"))
409
+ if isinstance(tag, CEnum) and tag.variant == "fixed_bucket" and width:
410
+ return {"kind": "fixed_bucket", "width_micros": width}
411
+ errors.raise_compile(
412
+ path,
413
+ errors.UNSUPPORTED_TYPE,
414
+ "cross-section grouping is neither exact_time nor a fixed bucket",
415
+ )
416
+ raise AssertionError("unreachable")
417
+
418
+
419
+ def _enum_attr(subtree: Node, name: str, /) -> str:
420
+ value = subtree.attr(name)
421
+ return value.variant if isinstance(value, CEnum) else ""
422
+
423
+
424
+ def _grouping_shape(subtree: Node, /) -> tuple[str, int] | None:
425
+ """Comparable grouping identity: the kind and, for buckets, the width."""
426
+
427
+ grouping = subtree.attr("grouping")
428
+ if isinstance(grouping, CEnum) and grouping.variant == "exact_time":
429
+ return ("exact_time", 0)
430
+ if isinstance(grouping, CMap):
431
+ tag = grouping.get("grouping")
432
+ width = _cint(grouping.get("width_micros"))
433
+ if isinstance(tag, CEnum) and tag.variant == "fixed_bucket" and width:
434
+ return ("fixed_bucket", width)
435
+ return None
436
+
437
+
438
+ # Cross-section planning validates every occurrence with stable,
439
+ # declaration-ordered error paths before emitting the frozen node shape.
440
+ # #lizard forgives
441
+ def _plan_cross_section(
442
+ output_name: str,
443
+ segment: _Segment,
444
+ path: str,
445
+ allowed_lateness_micros: int,
446
+ late_policy: str,
447
+ input_fields_override: tuple[Field, ...] | None,
448
+ /,
449
+ ) -> _CrossSectionPlan | None:
450
+ # #lizard forgives
451
+ occurrences: list[Node] = []
452
+ seen: set[str] = set()
453
+ ordered_trees = [tree for _, tree in segment.env]
454
+ ordered_trees += [
455
+ tree for tree in (segment.predicate, segment.post_predicate) if tree is not None
456
+ ]
457
+ for tree in ordered_trees:
458
+ for subtree in _find_cross_section(tree):
459
+ if subtree.digest not in seen:
460
+ seen.add(subtree.digest)
461
+ occurrences.append(subtree)
462
+ if not occurrences:
463
+ return None
464
+
465
+ input_fields = (
466
+ _schema_fields(segment.input_node.attr("schema"))
467
+ if input_fields_override is None
468
+ else input_fields_override
469
+ )
470
+ input_types = {field.name: field for field in input_fields}
471
+ entity_by = _cstr_seq(segment.input_node.attr("entity_by"))
472
+ sequence_by = _cstr_seq(segment.input_node.attr("sequence_by"))
473
+ event_time = _cstr(segment.input_node.attr("event_time"))
474
+ if not entity_by or not sequence_by or not event_time:
475
+ errors.raise_compile(
476
+ path,
477
+ errors.ORDERING_REQUIRED,
478
+ "cross-section primitives require declared entity_by,"
479
+ " event_time, and sequence_by ordering keys on the input table",
480
+ )
481
+
482
+ whole_feature = {
483
+ tree.digest: name
484
+ for name, tree in segment.env
485
+ if tree.op.name in _CROSS_SECTION_PRIMITIVES
486
+ }
487
+ used_names = set(input_types) | {name for name, _ in segment.env}
488
+ materialization = _plan_stateful_inputs(
489
+ _StatefulInputRequest(
490
+ output_name,
491
+ path,
492
+ "cs_input",
493
+ f"{output_name}__cf_cross_section_input",
494
+ "cross-section",
495
+ ),
496
+ input_fields,
497
+ used_names,
498
+ tuple((subtree.op.name, subtree.args[0]) for subtree in occurrences),
499
+ )
500
+ materializations = dict(materialization.names)
501
+ state_input_fields = materialization.input_fields
502
+ input_types = {field.name: field for field in state_input_fields}
503
+ used_names = set(materialization.used_names)
504
+ replacements: dict[str, str] = {}
505
+ declarations: list[dict[str, object]] = []
506
+ partition_columns: list[str] = []
507
+ derived_fields: list[Field] = []
508
+ for index, subtree in enumerate(occurrences):
509
+ kind = subtree.op.name
510
+ name = whole_feature.get(subtree.digest)
511
+ if name is None:
512
+ name = f"{output_name}__cf_cs_{index}"
513
+ if name in used_names:
514
+ errors.raise_compile(
515
+ f"{path}.{name}",
516
+ errors.DUPLICATE_NAME,
517
+ f"materialized cross-section column {name!r} collides with a"
518
+ " declared field",
519
+ )
520
+ used_names.add(name)
521
+ replacements[subtree.digest] = name
522
+ argument = subtree.args[0]
523
+ input_name = (
524
+ _cstr(argument.attr("name"))
525
+ if argument.op.name == "column_ref"
526
+ else materializations[argument.digest]
527
+ )
528
+ if input_name not in input_types:
529
+ errors.raise_compile(
530
+ f"{path}.{name}",
531
+ errors.SCHEMA_MISMATCH,
532
+ f"cross-section {kind} argument column {input_name!r} is not in"
533
+ " the input schema",
534
+ )
535
+ event_time_argument = subtree.args[1]
536
+ if event_time_argument.op.name != "column_ref":
537
+ errors.raise_compile(
538
+ f"{path}.{name}",
539
+ errors.UNSUPPORTED_TYPE,
540
+ "cross-section grouping event time must be an input column in"
541
+ " this release",
542
+ )
543
+ event_time_name = _cstr(event_time_argument.attr("name"))
544
+ if event_time_name != event_time:
545
+ errors.raise_compile(
546
+ f"{path}.{name}",
547
+ errors.SCHEMA_MISMATCH,
548
+ f"cross-section grouping event time {event_time_name!r} does not"
549
+ " match the declared input event time",
550
+ )
551
+ group_partitions: list[str] = []
552
+ for group_argument in subtree.args[2:]:
553
+ if group_argument.op.name != "column_ref":
554
+ errors.raise_compile(
555
+ f"{path}.{name}",
556
+ errors.UNSUPPORTED_TYPE,
557
+ "cross-section group columns must be input columns in this release",
558
+ )
559
+ group_name = _cstr(group_argument.attr("name"))
560
+ if group_name not in input_types:
561
+ errors.raise_compile(
562
+ f"{path}.{name}",
563
+ errors.SCHEMA_MISMATCH,
564
+ f"cross-section group column {group_name!r} is not in the"
565
+ " input schema",
566
+ )
567
+ group_partitions.append(group_name)
568
+ grouping_shape = _grouping_shape(subtree)
569
+ if index == 0:
570
+ partition_columns = group_partitions
571
+ declared_grouping = grouping_shape
572
+ elif (
573
+ partition_columns != group_partitions or declared_grouping != grouping_shape
574
+ ):
575
+ errors.raise_compile(
576
+ path,
577
+ errors.SCHEMA_MISMATCH,
578
+ "cross-section primitives in one output must share one"
579
+ " grouping declaration",
580
+ )
581
+ declaration: dict[str, object] = {
582
+ "kind": kind,
583
+ "primitive_version": 1,
584
+ "input": input_name,
585
+ "output": name,
586
+ }
587
+ if kind in _CROSS_SECTION_ORDERING:
588
+ declaration["direction"] = _enum_attr(subtree, "direction") or "ascending"
589
+ declaration["tie_method"] = _enum_attr(subtree, "tie_method") or "average"
590
+ declaration["null_placement"] = (
591
+ _enum_attr(subtree, "null_placement") or "exclude"
592
+ )
593
+ declaration["min_samples"] = _cint(subtree.attr("min_samples")) or 1
594
+ if kind == _CROSS_SECTION_DDOF:
595
+ declaration["ddof"] = _cint(subtree.attr("ddof")) or 0
596
+ if kind == "winsorize":
597
+ declaration["lower"] = _cnumber(subtree.attr("lower"))
598
+ declaration["upper"] = _cnumber(subtree.attr("upper"))
599
+ if kind in ("top", "bottom"):
600
+ declaration["count"] = _cint(subtree.attr("count"))
601
+ declaration["include_ties"] = _cbool(subtree.attr("include_ties"))
602
+ declarations.append(declaration)
603
+ output_type = (
604
+ "bool"
605
+ if kind in ("top", "bottom")
606
+ else input_types[input_name].data_type
607
+ if kind in ("winsorize", "mean_fill")
608
+ else "float64"
609
+ )
610
+ derived_fields.append(Field(name, output_type, nullable=True))
611
+
612
+ node_id = f"{output_name}__cf_cross_section"
613
+ node: dict[str, object] = {
614
+ "id": node_id,
615
+ "operator": {
616
+ "kind": "cross_section",
617
+ "spec": {
618
+ "configuration_version": 1,
619
+ "state_layout_version": 1,
620
+ "event_time": event_time,
621
+ "entity_by": list(entity_by),
622
+ "partition_by": list(partition_columns),
623
+ "sequence_by": list(sequence_by),
624
+ "grouping": _cross_section_grouping(
625
+ occurrences[0], f"{path}.{output_name}"
626
+ ),
627
+ "outputs": declarations,
628
+ "allowed_lateness_micros": allowed_lateness_micros,
629
+ "late_policy": (
630
+ {"kind": "error", "scope": "envelope"}
631
+ if late_policy == "error"
632
+ else {"kind": "drop", "metrics_version": 1}
633
+ ),
634
+ "value_policy": "nan_exclude_preserve_v1",
635
+ },
636
+ },
637
+ "input_ports": [
638
+ {
639
+ "name": "input",
640
+ "kind": "table",
641
+ "required": True,
642
+ "schema": [_field_json(field) for field in state_input_fields],
643
+ }
644
+ ],
645
+ "output_ports": [
646
+ {
647
+ "name": "output",
648
+ "kind": "table",
649
+ "required": True,
650
+ "schema": [
651
+ *(_field_json(field) for field in state_input_fields),
652
+ *(_field_json(field) for field in derived_fields),
653
+ ],
654
+ }
655
+ ],
656
+ }
657
+ env = tuple(
658
+ (name, _replace_materialized(tree, replacements)) for name, tree in segment.env
659
+ )
660
+ post_predicate = (
661
+ None
662
+ if segment.post_predicate is None
663
+ else _replace_materialized(segment.post_predicate, replacements)
664
+ )
665
+ return _CrossSectionPlan(
666
+ node_id,
667
+ node,
668
+ materialization.node_id,
669
+ materialization.node,
670
+ env,
671
+ post_predicate,
672
+ (*input_types, *(field.name for field in derived_fields)),
673
+ (*state_input_fields, *derived_fields),
674
+ materialization.names,
675
+ tuple(replacements.items()),
676
+ )
677
+
678
+
679
+ def _shared_materialized_name(
680
+ prefix: str, counter: int, reserved: set[str], /
681
+ ) -> tuple[str, int]:
682
+ while True:
683
+ name = f"__cf_shared_{prefix}_{counter}"
684
+ counter += 1
685
+ if name not in reserved:
686
+ reserved.add(name)
687
+ return name, counter
688
+
689
+
690
+ def _plan_outputs(plan: _RollingPlan | _CrossSectionPlan, /) -> list[dict[str, object]]:
691
+ return plan.node["operator"]["spec"]["outputs"] # type: ignore[index,return-value]
692
+
693
+
694
+ def _required_state_plan[StatePlanT: (_RollingPlan, _CrossSectionPlan)](
695
+ plans: dict[str, StatePlanT | None], output_name: str, /
696
+ ) -> StatePlanT:
697
+ plan = plans[output_name]
698
+ if plan is None:
699
+ raise RuntimeError(f"missing shared-state plan for output {output_name!r}")
700
+ return plan
701
+
702
+
703
+ def _merge_state_outputs[StatePlanT: (_RollingPlan, _CrossSectionPlan)](
704
+ members: list[str],
705
+ plans: dict[str, StatePlanT | None],
706
+ prefix: str,
707
+ /,
708
+ ) -> tuple[StatePlanT, dict[str, str], list[dict[str, object]], tuple[Field, ...]]:
709
+ first = _required_state_plan(plans, members[0])
710
+ base_count = len(first.output_fields) - len(first.replacements)
711
+ base_fields = first.output_fields[:base_count]
712
+ reserved = {
713
+ field.name
714
+ for output_name in members
715
+ for field in _required_state_plan(plans, output_name).output_fields
716
+ }
717
+ replacements: dict[str, str] = {}
718
+ declarations: list[dict[str, object]] = []
719
+ derived_fields: list[Field] = []
720
+ counter = 0
721
+ for output_name in members:
722
+ plan = _required_state_plan(plans, output_name)
723
+ declarations_by_name = {item["output"]: item for item in _plan_outputs(plan)}
724
+ fields_by_name = {field.name: field for field in plan.output_fields}
725
+ for digest, old_name in plan.replacements:
726
+ if digest in replacements:
727
+ continue
728
+ new_name, counter = _shared_materialized_name(prefix, counter, reserved)
729
+ replacements[digest] = new_name
730
+ declarations.append({**declarations_by_name[old_name], "output": new_name})
731
+ field = fields_by_name[old_name]
732
+ derived_fields.append(
733
+ Field(new_name, field.data_type, nullable=field.nullable)
734
+ )
735
+ return first, replacements, declarations, (*base_fields, *derived_fields)
736
+
737
+
738
+ def _shared_state_node(
739
+ plan: _RollingPlan | _CrossSectionPlan,
740
+ node_id: str,
741
+ declarations: list[dict[str, object]],
742
+ output_fields: tuple[Field, ...],
743
+ /,
744
+ ) -> dict[str, object]:
745
+ operator = plan.node["operator"]
746
+ spec = operator["spec"] # type: ignore[index]
747
+ return {
748
+ **plan.node,
749
+ "id": node_id,
750
+ "operator": {
751
+ **operator, # type: ignore[arg-type]
752
+ "spec": {**spec, "outputs": declarations}, # type: ignore[arg-type]
753
+ },
754
+ "output_ports": [
755
+ {
756
+ "name": "output",
757
+ "kind": "table",
758
+ "required": True,
759
+ "schema": [_field_json(field) for field in output_fields],
760
+ }
761
+ ],
762
+ }
763
+
764
+
765
+ def _unique_shared_node_id(stem: str, reserved: set[str], /) -> str:
766
+ candidate = stem
767
+ counter = 1
768
+ while candidate in reserved:
769
+ candidate = f"{stem}_{counter}"
770
+ counter += 1
771
+ reserved.add(candidate)
772
+ return candidate
773
+
774
+
775
+ def _reserved_state_node_ids[StatePlanT: (_RollingPlan, _CrossSectionPlan)](
776
+ segments: list[tuple[str, _Segment]],
777
+ plans: dict[str, StatePlanT | None],
778
+ /,
779
+ ) -> set[str]:
780
+ reserved = {output_name for output_name, _ in segments}
781
+ reserved.update(_cstr(segment.input_node.attr("name")) for _, segment in segments)
782
+ reserved.update(plan.node_id for plan in plans.values() if plan is not None)
783
+ return reserved
784
+
785
+
786
+ def _shared_group_plans[StatePlanT: (_RollingPlan, _CrossSectionPlan)](
787
+ members: list[str],
788
+ by_name: dict[str, _Segment],
789
+ first: StatePlanT,
790
+ replacements: dict[str, str],
791
+ node_id: str,
792
+ node: dict[str, object],
793
+ output_fields: tuple[Field, ...],
794
+ /,
795
+ ) -> dict[str, StatePlanT]:
796
+ shared: dict[str, StatePlanT] = {}
797
+ for output_name in members:
798
+ segment = by_name[output_name]
799
+ post_predicate = segment.post_predicate
800
+ if post_predicate is not None:
801
+ post_predicate = _replace_materialized(post_predicate, replacements)
802
+ shared[output_name] = replace(
803
+ first,
804
+ node_id=node_id,
805
+ node=node,
806
+ env=tuple(
807
+ (name, _replace_materialized(tree, replacements))
808
+ for name, tree in segment.env
809
+ ),
810
+ post_predicate=post_predicate,
811
+ input_field_names=tuple(field.name for field in output_fields),
812
+ output_fields=output_fields,
813
+ replacements=tuple(replacements.items()),
814
+ )
815
+ return shared
816
+
817
+
818
+ def _share_state_groups[StatePlanT: (_RollingPlan, _CrossSectionPlan)](
819
+ groups: list[list[str]],
820
+ segments: list[tuple[str, _Segment]],
821
+ plans: dict[str, StatePlanT | None],
822
+ prefix: str,
823
+ suffix: str,
824
+ /,
825
+ ) -> dict[str, StatePlanT | None]:
826
+ shared = dict(plans)
827
+ by_name = dict(segments)
828
+ reserved_ids = _reserved_state_node_ids(segments, plans)
829
+ for members in groups:
830
+ if len(members) < 2:
831
+ continue
832
+ first, replacements, declarations, output_fields = _merge_state_outputs(
833
+ members, plans, prefix
834
+ )
835
+ node_id = _unique_shared_node_id(
836
+ f"{members[0]}__cf_shared_{suffix}", reserved_ids
837
+ )
838
+ node = _shared_state_node(first, node_id, declarations, output_fields)
839
+ shared.update(
840
+ _shared_group_plans(
841
+ members,
842
+ by_name,
843
+ first,
844
+ replacements,
845
+ node_id,
846
+ node,
847
+ output_fields,
848
+ )
849
+ )
850
+ return shared
851
+
852
+
853
+ def _share_single_rolling_stages(
854
+ segments: list[tuple[str, _Segment]],
855
+ plans: dict[str, _RollingPlan | None],
856
+ /,
857
+ ) -> dict[str, _RollingPlan | None]:
858
+ groups: dict[tuple[str, str | None, str | None], list[str]] = {}
859
+ for output_name, segment in segments:
860
+ plan = plans[output_name]
861
+ if plan is not None:
862
+ predicate = None if segment.predicate is None else segment.predicate.digest
863
+ materialization = (
864
+ None
865
+ if plan.materialization_node is None
866
+ else _canonical(plan.materialization_node)
867
+ )
868
+ groups.setdefault(
869
+ (segment.input_node.digest, predicate, materialization), []
870
+ ).append(output_name)
871
+ return _share_state_groups(
872
+ list(groups.values()), segments, plans, "roll", "rolling"
873
+ )
874
+
875
+
876
+ def _share_rolling_plans(
877
+ segments: list[tuple[str, _Segment]],
878
+ plans: dict[str, _RollingPipeline | None],
879
+ /,
880
+ ) -> dict[str, _RollingPipeline | None]:
881
+ """Preserve existing cross-output sharing for one-stage pipelines."""
882
+
883
+ stages = {
884
+ output_name: (
885
+ None
886
+ if pipeline is None or len(pipeline.stages) != 1
887
+ else pipeline.stages[0]
888
+ )
889
+ for output_name, pipeline in plans.items()
890
+ }
891
+ shared_stages = _share_single_rolling_stages(segments, stages)
892
+ shared = dict(plans)
893
+ for output_name, stage in shared_stages.items():
894
+ pipeline = plans[output_name]
895
+ if pipeline is None or stage is None:
896
+ continue
897
+ shared[output_name] = _RollingPipeline(
898
+ (stage,),
899
+ stage.env,
900
+ stage.post_predicate,
901
+ stage.input_field_names,
902
+ stage.output_fields,
903
+ )
904
+ return _share_identical_multi_stage_pipelines(segments, shared)
905
+
906
+
907
+ def _rolling_pipeline_identity(
908
+ segment: _Segment, /
909
+ ) -> tuple[str, str | None, tuple[str, ...]]:
910
+ predicate = None if segment.predicate is None else segment.predicate.digest
911
+ seen: set[str] = set()
912
+ digests: list[str] = []
913
+ trees = [tree for _, tree in segment.env]
914
+ trees += [tree for tree in (segment.post_predicate,) if tree is not None]
915
+ for tree in trees:
916
+ for subtree in _find_rolling(tree):
917
+ if subtree.digest not in seen:
918
+ seen.add(subtree.digest)
919
+ digests.append(subtree.digest)
920
+ return segment.input_node.digest, predicate, tuple(digests)
921
+
922
+
923
+ def _shared_pipeline_stages(
924
+ first_name: str,
925
+ pipeline: _RollingPipeline,
926
+ reserved_ids: set[str],
927
+ /,
928
+ ) -> tuple[_RollingPlan, ...]:
929
+ stages: list[_RollingPlan] = []
930
+ for index, stage in enumerate(pipeline.stages, start=1):
931
+ node_id = _unique_shared_node_id(
932
+ f"{first_name}__cf_shared_rolling_{index}", reserved_ids
933
+ )
934
+ materialization_id = None
935
+ materialization = None
936
+ if stage.materialization_node is not None:
937
+ materialization_id = _unique_shared_node_id(
938
+ f"{first_name}__cf_shared_rolling_input_{index}", reserved_ids
939
+ )
940
+ materialization = {
941
+ **stage.materialization_node,
942
+ "id": materialization_id,
943
+ }
944
+ stages.append(
945
+ replace(
946
+ stage,
947
+ node_id=node_id,
948
+ node={**stage.node, "id": node_id},
949
+ materialization_node_id=materialization_id,
950
+ materialization_node=materialization,
951
+ )
952
+ )
953
+ return tuple(stages)
954
+
955
+
956
+ def _rewrite_pipeline_environment(
957
+ segment: _Segment,
958
+ stages: tuple[_RollingPlan, ...],
959
+ /,
960
+ ) -> tuple[tuple[tuple[str, Node], ...], Node | None]:
961
+ env = segment.env
962
+ post_predicate = segment.post_predicate
963
+ for stage in stages:
964
+ replacements = dict(stage.replacements)
965
+ env = tuple(
966
+ (name, _replace_materialized(tree, replacements)) for name, tree in env
967
+ )
968
+ if post_predicate is not None:
969
+ post_predicate = _replace_materialized(post_predicate, replacements)
970
+ return env, post_predicate
971
+
972
+
973
+ def _multi_stage_pipeline_groups(
974
+ segments: list[tuple[str, _Segment]],
975
+ plans: dict[str, _RollingPipeline | None],
976
+ /,
977
+ ) -> tuple[tuple[str, ...], ...]:
978
+ groups: dict[tuple[str, str | None, tuple[str, ...]], list[str]] = {}
979
+ for output_name, segment in segments:
980
+ pipeline = plans[output_name]
981
+ if pipeline is not None and len(pipeline.stages) > 1:
982
+ groups.setdefault(_rolling_pipeline_identity(segment), []).append(
983
+ output_name
984
+ )
985
+ return tuple(tuple(members) for members in groups.values() if len(members) > 1)
986
+
987
+
988
+ def _reserved_rolling_pipeline_ids(
989
+ segments: list[tuple[str, _Segment]],
990
+ plans: dict[str, _RollingPipeline | None],
991
+ /,
992
+ ) -> set[str]:
993
+ reserved_ids = {output_name for output_name, _ in segments}
994
+ for pipeline in plans.values():
995
+ if pipeline is None:
996
+ continue
997
+ for stage in pipeline.stages:
998
+ reserved_ids.add(stage.node_id)
999
+ if stage.materialization_node_id is not None:
1000
+ reserved_ids.add(stage.materialization_node_id)
1001
+ return reserved_ids
1002
+
1003
+
1004
+ def _shared_multi_stage_group(
1005
+ members: tuple[str, ...],
1006
+ segments: dict[str, _Segment],
1007
+ plans: dict[str, _RollingPipeline | None],
1008
+ reserved_ids: set[str],
1009
+ /,
1010
+ ) -> dict[str, _RollingPipeline]:
1011
+ first = plans[members[0]]
1012
+ if first is None:
1013
+ raise RuntimeError("missing multi-stage rolling pipeline")
1014
+ stages = _shared_pipeline_stages(members[0], first, reserved_ids)
1015
+ shared: dict[str, _RollingPipeline] = {}
1016
+ for output_name in members:
1017
+ env, post_predicate = _rewrite_pipeline_environment(
1018
+ segments[output_name], stages
1019
+ )
1020
+ shared[output_name] = _RollingPipeline(
1021
+ stages,
1022
+ env,
1023
+ post_predicate,
1024
+ first.input_field_names,
1025
+ first.output_fields,
1026
+ )
1027
+ return shared
1028
+
1029
+
1030
+ def _share_identical_multi_stage_pipelines(
1031
+ segments: list[tuple[str, _Segment]],
1032
+ plans: dict[str, _RollingPipeline | None],
1033
+ /,
1034
+ ) -> dict[str, _RollingPipeline | None]:
1035
+ groups = _multi_stage_pipeline_groups(segments, plans)
1036
+ reserved_ids = _reserved_rolling_pipeline_ids(segments, plans)
1037
+ segments_by_name = dict(segments)
1038
+ shared = dict(plans)
1039
+ for members in groups:
1040
+ shared.update(
1041
+ _shared_multi_stage_group(
1042
+ members,
1043
+ segments_by_name,
1044
+ plans,
1045
+ reserved_ids,
1046
+ )
1047
+ )
1048
+ return shared
1049
+
1050
+
1051
+ def _cross_section_group_identity(plan: _CrossSectionPlan, /) -> str:
1052
+ spec = plan.node["operator"]["spec"] # type: ignore[index]
1053
+ return _canonical({key: value for key, value in spec.items() if key != "outputs"})
1054
+
1055
+
1056
+ def _materialized_select_expression(plan: _CrossSectionPlan, name: str, /) -> str:
1057
+ node = plan.materialization_node
1058
+ if node is None:
1059
+ raise RuntimeError(f"missing cross-section materialization for {name!r}")
1060
+ suffix = f" AS {_quote_identifier(name)}"
1061
+ selects = node["operator"]["select"] # type: ignore[index]
1062
+ for item in selects:
1063
+ if item.endswith(suffix):
1064
+ return item[: -len(suffix)]
1065
+ raise RuntimeError(f"missing materialized select for {name!r}")
1066
+
1067
+
1068
+ @dataclass(frozen=True, slots=True)
1069
+ class _SharedCrossInputs:
1070
+ base_fields: tuple[Field, ...]
1071
+ fields: tuple[Field, ...]
1072
+ expressions: tuple[str, ...]
1073
+ names: tuple[tuple[str, str], ...]
1074
+
1075
+ @property
1076
+ def input_fields(self) -> tuple[Field, ...]:
1077
+ return (*self.base_fields, *self.fields)
1078
+
1079
+
1080
+ def _collect_shared_cross_inputs(
1081
+ members: list[str],
1082
+ plans: dict[str, _CrossSectionPlan | None],
1083
+ /,
1084
+ ) -> _SharedCrossInputs:
1085
+ first = _required_state_plan(plans, members[0])
1086
+ base_count = (
1087
+ len(first.output_fields) - len(first.replacements) - len(first.materializations)
1088
+ )
1089
+ reserved = {
1090
+ field.name
1091
+ for output_name in members
1092
+ for field in _required_state_plan(plans, output_name).output_fields
1093
+ }
1094
+ names: dict[str, str] = {}
1095
+ fields: list[Field] = []
1096
+ expressions: list[str] = []
1097
+ counter = 0
1098
+ for output_name in members:
1099
+ plan = _required_state_plan(plans, output_name)
1100
+ fields_by_name = {field.name: field for field in plan.output_fields}
1101
+ for digest, old_name in plan.materializations:
1102
+ if digest in names:
1103
+ continue
1104
+ name, counter = _shared_materialized_name("cs_input", counter, reserved)
1105
+ names[digest] = name
1106
+ field = fields_by_name[old_name]
1107
+ fields.append(Field(name, field.data_type, nullable=field.nullable))
1108
+ expressions.append(_materialized_select_expression(plan, old_name))
1109
+ return _SharedCrossInputs(
1110
+ first.output_fields[:base_count],
1111
+ tuple(fields),
1112
+ tuple(expressions),
1113
+ tuple(names.items()),
1114
+ )
1115
+
1116
+
1117
+ def _shared_cross_materialization(
1118
+ shared: _SharedCrossInputs, node_id: str, /
1119
+ ) -> dict[str, object] | None:
1120
+ if not shared.fields:
1121
+ return None
1122
+ return _expression_node(
1123
+ node_id,
1124
+ [
1125
+ *(_quote_identifier(field.name) for field in shared.base_fields),
1126
+ *(
1127
+ f"{expression} AS {_quote_identifier(field.name)}"
1128
+ for expression, field in zip(
1129
+ shared.expressions, shared.fields, strict=True
1130
+ )
1131
+ ),
1132
+ ],
1133
+ None,
1134
+ shared.base_fields,
1135
+ shared.input_fields,
1136
+ )
1137
+
1138
+
1139
+ def _aligned_cross_section_plan(
1140
+ plan: _CrossSectionPlan,
1141
+ shared: _SharedCrossInputs,
1142
+ node_id: str,
1143
+ materialization: dict[str, object] | None,
1144
+ /,
1145
+ ) -> _CrossSectionPlan:
1146
+ names = dict(shared.names)
1147
+ old_to_new = {old_name: names[digest] for digest, old_name in plan.materializations}
1148
+ outputs = [
1149
+ {**item, "input": old_to_new.get(item["input"], item["input"])}
1150
+ for item in _plan_outputs(plan)
1151
+ ]
1152
+ derived_fields = plan.output_fields[-len(plan.replacements) :]
1153
+ output_fields = (*shared.input_fields, *derived_fields)
1154
+ return replace(
1155
+ plan,
1156
+ node={
1157
+ **plan.node,
1158
+ "operator": {
1159
+ **plan.node["operator"], # type: ignore[dict-item]
1160
+ "spec": {
1161
+ **plan.node["operator"]["spec"], # type: ignore[index]
1162
+ "outputs": outputs,
1163
+ },
1164
+ },
1165
+ "input_ports": [
1166
+ {
1167
+ "name": "input",
1168
+ "kind": "table",
1169
+ "required": True,
1170
+ "schema": [_field_json(field) for field in shared.input_fields],
1171
+ }
1172
+ ],
1173
+ "output_ports": [
1174
+ {
1175
+ "name": "output",
1176
+ "kind": "table",
1177
+ "required": True,
1178
+ "schema": [_field_json(field) for field in output_fields],
1179
+ }
1180
+ ],
1181
+ },
1182
+ materialization_node_id=node_id if shared.fields else None,
1183
+ materialization_node=materialization,
1184
+ input_field_names=tuple(field.name for field in output_fields),
1185
+ output_fields=output_fields,
1186
+ materializations=shared.names,
1187
+ )
1188
+
1189
+
1190
+ def _combined_cross_section_inputs(
1191
+ members: list[str],
1192
+ plans: dict[str, _CrossSectionPlan | None],
1193
+ node_id: str,
1194
+ /,
1195
+ ) -> dict[str, _CrossSectionPlan]:
1196
+ shared = _collect_shared_cross_inputs(members, plans)
1197
+ materialization = _shared_cross_materialization(shared, node_id)
1198
+ return {
1199
+ output_name: _aligned_cross_section_plan(
1200
+ _required_state_plan(plans, output_name),
1201
+ shared,
1202
+ node_id,
1203
+ materialization,
1204
+ )
1205
+ for output_name in members
1206
+ }
1207
+
1208
+
1209
+ def _share_cross_section_plans(
1210
+ segments: list[tuple[str, _Segment]],
1211
+ upstream_ids: dict[str, str | None],
1212
+ plans: dict[str, _CrossSectionPlan | None],
1213
+ /,
1214
+ ) -> dict[str, _CrossSectionPlan | None]:
1215
+ groups: dict[tuple[str, str | None, str], list[str]] = {}
1216
+ for output_name, segment in segments:
1217
+ plan = plans[output_name]
1218
+ if plan is not None:
1219
+ upstream = upstream_ids[output_name] or segment.input_node.digest
1220
+ predicate = None if segment.predicate is None else segment.predicate.digest
1221
+ groups.setdefault(
1222
+ (upstream, predicate, _cross_section_group_identity(plan)), []
1223
+ ).append(output_name)
1224
+ aligned = dict(plans)
1225
+ reserved_ids = _reserved_state_node_ids(segments, plans)
1226
+ for group in groups.values():
1227
+ if len(group) < 2:
1228
+ continue
1229
+ materialization_id = _unique_shared_node_id(
1230
+ f"{group[0]}__cf_shared_cross_section_input", reserved_ids
1231
+ )
1232
+ aligned.update(_combined_cross_section_inputs(group, plans, materialization_id))
1233
+ return _share_state_groups(
1234
+ list(groups.values()), segments, aligned, "cs", "cross_section"
1235
+ )
1236
+
1237
+
1238
+ def _check_declared_inputs(program: Program, /) -> None:
1239
+ for value in program.inputs:
1240
+ node = value._node
1241
+ if node.op.name == "parameter":
1242
+ errors.raise_compile(
1243
+ f"static_inputs.{_cstr(node.attr('name'))}",
1244
+ errors.UNKNOWN_PRIMITIVE_VERSION,
1245
+ "static parameters are not supported by the row-local lowerer",
1246
+ )
1247
+
1248
+
1249
+ @dataclass(frozen=True, slots=True)
1250
+ class _MatrixExpression:
1251
+ backend: str
1252
+ columns: tuple[str, ...]
1253
+ source_digests: frozenset[str]
1254
+ parameter: Node | None
1255
+ weights_count: int
1256
+ matmul_count: int
1257
+ matmul_rhs_is_weights: bool
1258
+ tree: dict[str, object]
1259
+
1260
+
1261
+ @dataclass(frozen=True, slots=True)
1262
+ class _LoweringValue:
1263
+ _node: Node
1264
+
1265
+
1266
+ @dataclass(frozen=True, slots=True)
1267
+ class _LoweringProgram:
1268
+ name: str
1269
+ inputs: tuple[_LoweringValue, ...]
1270
+ outputs: tuple[tuple[str, _LoweringValue], ...]