Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Prev Previous commit
Next Next commit
update doc
Signed-off-by: HaoXuAI <sduxuhao@gmail.com>
  • Loading branch information
haoxu0 committed Apr 17, 2025
commit f89ebb16ac61e8f32508315541228685dc852622
37 changes: 25 additions & 12 deletions sdk/python/feast/infra/compute_engines/local/feature_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,9 +19,9 @@

class LocalFeatureBuilder(FeatureBuilder):
def __init__(
self,
task: Union[MaterializationTask, HistoricalRetrievalTask],
backend: DataFrameBackend,
self,
task: Union[MaterializationTask, HistoricalRetrievalTask],
backend: DataFrameBackend,
):
super().__init__(task)
self.backend = backend
Expand All @@ -31,13 +31,15 @@ def build_source_node(self):
self.nodes.append(node)
return node

def build_join_node(self, input_node):
def build_join_node(self,
input_node):
node = LocalJoinNode("join", self.backend)
node.add_input(input_node)
self.nodes.append(node)
return node

def build_filter_node(self, input_node):
def build_filter_node(self,
input_node):
filter_expr = None
if hasattr(self.feature_view, "filter"):
filter_expr = self.feature_view.filter
Expand All @@ -47,45 +49,56 @@ def build_filter_node(self, input_node):
self.nodes.append(node)
return node

def build_aggregation_node(self, input_node):
agg_specs = self.feature_view.aggregations
@staticmethod
def _get_aggregate_operations(agg_specs):
agg_ops = {}
for agg in agg_specs:
if agg.time_window is not None:
raise ValueError(
"Time window aggregation is not supported in local compute engine. Please use a different compute engine."
"Time window aggregation is not supported in local compute engine. Please use a different compute "
"engine."
)
alias = f"{agg.function}_{agg.column}"
agg_ops[alias] = (agg.function, agg.column)
return agg_ops

def build_aggregation_node(self,
input_node):
agg_specs = self.feature_view.aggregations
agg_ops = self._get_aggregate_operations(agg_specs)
group_by_keys = self.feature_view.entities
node = LocalAggregationNode("agg", self.backend, group_by_keys, agg_ops)
node.add_input(input_node)
self.nodes.append(node)
return node

def build_dedup_node(self, input_node):
def build_dedup_node(self,
input_node):
node = LocalDedupNode("dedup", self.backend)
node.add_input(input_node)
self.nodes.append(node)
return node

def build_transformation_node(self, input_node):
def build_transformation_node(self,
input_node):
node = LocalTransformationNode(
"transform", self.feature_view.feature_transformation, self.backend
)
node.add_input(input_node)
self.nodes.append(node)
return node

def build_validation_node(self, input_node):
def build_validation_node(self,
input_node):
node = LocalValidationNode(
"validate", self.feature_view.validation_config, self.backend
)
node.add_input(input_node)
self.nodes.append(node)
return node

def build_output_nodes(self, input_node):
def build_output_nodes(self,
input_node):
node = LocalOutputNode("output")
node.add_input(input_node)
self.nodes.append(node)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -55,8 +55,9 @@ def build_filter_node(self, input_node):
filter_expr = None
if hasattr(self.feature_view, "filter"):
filter_expr = self.feature_view.filter
ttl = self.feature_view.ttl
node = SparkFilterNode(
"filter", self.spark_session, self.feature_view, filter_expr
"filter", self.spark_session, ttl, filter_expr
)
node.add_input(input_node)
self.nodes.append(node)
Expand Down
27 changes: 6 additions & 21 deletions sdk/python/feast/infra/compute_engines/spark/node.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
from dataclasses import dataclass
from datetime import datetime
from typing import Dict, List, Optional, Union, cast
from datetime import timedelta
from typing import List, Optional, Union, cast

from pyspark.sql import DataFrame, SparkSession, Window
from pyspark.sql import functions as F
Expand Down Expand Up @@ -50,20 +49,6 @@ def rename_entity_ts_column(
return entity_df


@dataclass
class SparkJoinContext:
name: str # feature view name or alias
join_keys: List[str]
feature_columns: List[str]
timestamp_field: str
created_timestamp_column: Optional[str]
ttl_seconds: Optional[int]
min_event_timestamp: Optional[datetime]
max_event_timestamp: Optional[datetime]
field_mapping: Dict[str, str] # original_column_name -> renamed_column
full_feature_names: bool = False # apply feature view name prefix


class SparkMaterializationReadNode(DAGNode):
def __init__(
self, name: str, task: Union[MaterializationTask, HistoricalRetrievalTask]
Expand Down Expand Up @@ -266,12 +251,12 @@ def __init__(
self,
name: str,
spark_session: SparkSession,
feature_view: Union[BatchFeatureView, StreamFeatureView],
ttl: Optional[timedelta] = None,
filter_condition: Optional[str] = None,
):
super().__init__(name)
self.spark_session = spark_session
self.feature_view = feature_view
self.ttl = ttl
self.filter_condition = filter_condition

def execute(self, context: ExecutionContext) -> DAGValue:
Expand All @@ -288,8 +273,8 @@ def execute(self, context: ExecutionContext) -> DAGValue:
filtered_df = filtered_df.filter(F.col(ts_col) <= F.col(ENTITY_TS_ALIAS))

# Optional TTL filter: feature.ts >= entity.event_timestamp - ttl
if self.feature_view.ttl:
ttl_seconds = int(self.feature_view.ttl.total_seconds())
if self.ttl:
ttl_seconds = int(self.ttl.total_seconds())
lower_bound = F.col(ENTITY_TS_ALIAS) - F.expr(
f"INTERVAL {ttl_seconds} seconds"
)
Expand Down