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,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], ...]
|