Repository navigation
Expand file tree
/
Copy pathjob.py
More file actions
298 lines (260 loc) · 11.1 KB
/
Copy pathjob.py
File metadata and controls
298 lines (260 loc) · 11.1 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
import logging
import uuid
from dataclasses import dataclass
from typing import List, Optional
import pandas as pd
import pyarrow as pa
from ray.data import Dataset
from feast import OnDemandFeatureView
from feast.dqm.errors import ValidationFailed
from feast.errors import SavedDatasetLocationAlreadyExists
from feast.infra.common.materialization_job import (
MaterializationJob,
MaterializationJobStatus,
)
from feast.infra.compute_engines.dag.context import ExecutionContext
from feast.infra.compute_engines.dag.model import DAGFormat
from feast.infra.compute_engines.dag.plan import ExecutionPlan
from feast.infra.compute_engines.dag.value import DAGValue
from feast.infra.offline_stores.file_source import SavedDatasetFileStorage
from feast.infra.offline_stores.offline_store import RetrievalJob, RetrievalMetadata
from feast.infra.ray_initializer import get_ray_wrapper
from feast.repo_config import RepoConfig
from feast.saved_dataset import SavedDatasetStorage
logger = logging.getLogger(__name__)
class RayDAGRetrievalJob(RetrievalJob):
"""
Ray-based retrieval job that executes a DAG plan to retrieve historical features.
"""
def __init__(
self,
plan: Optional[ExecutionPlan],
context: Optional[ExecutionContext],
config: RepoConfig,
full_feature_names: bool,
on_demand_feature_views: Optional[List[OnDemandFeatureView]] = None,
feature_refs: Optional[List[str]] = None,
metadata: Optional[RetrievalMetadata] = None,
error: Optional[BaseException] = None,
):
super().__init__()
self._plan = plan
self._context = context
self._config = config
self._full_feature_names = full_feature_names
self._on_demand_feature_views = on_demand_feature_views or []
self._feature_refs = feature_refs or []
self._metadata = metadata
self._error = error
self._result_dataset: Optional[Dataset] = None
self._result_df: Optional[pd.DataFrame] = None
self._result_arrow: Optional[pa.Table] = None
def error(self) -> Optional[BaseException]:
"""Return any error that occurred during job execution."""
return self._error
def _ensure_executed(self) -> DAGValue:
"""Ensure the execution plan has been executed."""
if self._result_dataset is None and self._plan and self._context:
try:
result = self._plan.execute(self._context)
if hasattr(result, "data") and isinstance(result.data, Dataset):
self._result_dataset = result.data
else:
# If result is not a Ray Dataset, convert it
ray_wrapper = get_ray_wrapper()
if isinstance(result.data, pd.DataFrame):
self._result_dataset = ray_wrapper.from_pandas(result.data)
elif isinstance(result.data, pa.Table):
self._result_dataset = ray_wrapper.from_arrow(result.data)
else:
raise ValueError(
f"Unsupported result type: {type(result.data)}"
)
return result
except Exception as e:
self._error = e
logger.error(f"Ray DAG execution failed: {e}")
raise
elif self._result_dataset is None:
raise ValueError("No execution plan available or execution failed")
# Return a mock DAGValue for compatibility
return DAGValue(data=self._result_dataset, format=DAGFormat.RAY)
def to_ray_dataset(self) -> Dataset:
"""Get the result as a Ray Dataset."""
self._ensure_executed()
assert self._result_dataset is not None, (
"Dataset should not be None after execution"
)
return self._result_dataset
def to_df(
self,
validation_reference=None,
timeout: Optional[int] = None,
) -> pd.DataFrame:
"""Convert the result to a pandas DataFrame."""
if self._result_df is None:
if self.on_demand_feature_views:
# Use parent implementation for ODFV processing
logger.info(
f"Processing {len(self.on_demand_feature_views)} on-demand feature views"
)
self._result_df = super().to_df(
validation_reference=validation_reference, timeout=timeout
)
else:
# Direct conversion from Ray Dataset
self._ensure_executed()
assert self._result_dataset is not None, (
"Dataset should not be None after execution"
)
self._result_df = self._result_dataset.to_pandas()
# Handle validation if provided
if validation_reference:
try:
validation_result = validation_reference.profile.validate(
self._result_df
)
if not validation_result.is_success:
raise ValidationFailed(validation_result)
except ImportError:
logger.warning("DQM profiler not available, skipping validation")
except Exception as e:
logger.error(f"Validation failed: {e}")
raise ValueError(f"Data validation failed: {e}")
return self._result_df
def to_arrow(
self,
validation_reference=None,
timeout: Optional[int] = None,
) -> pa.Table:
"""Convert the result to an Arrow Table."""
if self._result_arrow is None:
if self.on_demand_feature_views:
# Use parent implementation for ODFV processing
self._result_arrow = super().to_arrow(
validation_reference=validation_reference, timeout=timeout
)
else:
# Direct conversion from Ray Dataset
self._ensure_executed()
assert self._result_dataset is not None, (
"Dataset should not be None after execution"
)
self._result_arrow = self._result_dataset.to_pandas().to_arrow()
# Handle validation if provided
if validation_reference:
try:
df = self._result_arrow.to_pandas()
validation_result = validation_reference.profile.validate(df)
if not validation_result.is_success:
raise ValidationFailed(validation_result)
except ImportError:
logger.warning("DQM profiler not available, skipping validation")
except Exception as e:
logger.error(f"Validation failed: {e}")
raise ValueError(f"Data validation failed: {e}")
return self._result_arrow
def to_remote_storage(self) -> list[str]:
"""Write the result to remote storage."""
if not self._config.batch_engine.staging_location:
raise ValueError("Staging location must be set for remote storage")
try:
self._ensure_executed()
assert self._result_dataset is not None, (
"Dataset should not be None after execution"
)
output_uri = (
f"{self._config.batch_engine.staging_location}/{str(uuid.uuid4())}"
)
self._result_dataset.write_parquet(output_uri)
logger.debug(f"Wrote result to {output_uri}")
return [output_uri]
except Exception as e:
raise RuntimeError(f"Failed to write to remote storage: {e}")
def persist(
self,
storage: SavedDatasetStorage,
allow_overwrite: bool = False,
timeout: Optional[int] = None,
) -> str:
"""Persist the result to the specified storage."""
if not isinstance(storage, SavedDatasetFileStorage):
raise ValueError(
f"Ray compute engine only supports SavedDatasetFileStorage, got {type(storage)}"
)
destination_path = storage.file_options.uri
# Check if destination already exists
if not destination_path.startswith(("s3://", "gs://", "hdfs://")):
import os
if not allow_overwrite and os.path.exists(destination_path):
raise SavedDatasetLocationAlreadyExists(location=destination_path)
os.makedirs(os.path.dirname(destination_path), exist_ok=True)
try:
self._ensure_executed()
assert self._result_dataset is not None, (
"Dataset should not be None after execution"
)
self._result_dataset.write_parquet(destination_path)
return destination_path
except Exception as e:
raise RuntimeError(f"Failed to persist dataset to {destination_path}: {e}")
def to_sql(self) -> str:
"""Generate SQL representation of the execution plan."""
if self._plan and self._context:
return self._plan.to_sql(self._context)
raise NotImplementedError("SQL generation not available without execution plan")
@property
def full_feature_names(self) -> bool:
return self._full_feature_names
@property
def on_demand_feature_views(self) -> List[OnDemandFeatureView]:
return self._on_demand_feature_views
@property
def metadata(self) -> Optional[RetrievalMetadata]:
return self._metadata
def _to_df_internal(self, timeout: Optional[int] = None) -> pd.DataFrame:
"""Internal method to get DataFrame (used by parent class)."""
self._ensure_executed()
assert self._result_dataset is not None, (
"Dataset should not be None after execution"
)
return self._result_dataset.to_pandas()
def _to_arrow_internal(self, timeout: Optional[int] = None) -> pa.Table:
"""Internal method to get Arrow Table (used by parent class)."""
self._ensure_executed()
assert self._result_dataset is not None, (
"Dataset should not be None after execution"
)
return self._result_dataset.to_pandas().to_arrow()
@dataclass
class RayMaterializationJob(MaterializationJob):
"""
Ray-based materialization job that tracks the status of feature materialization.
"""
def __init__(
self,
job_id: str,
status: MaterializationJobStatus,
result: Optional[DAGValue] = None,
error: Optional[BaseException] = None,
):
super().__init__()
self._job_id = job_id
self._status = status
self._result = result
self._error = error
def job_id(self) -> str:
return self._job_id
def status(self) -> MaterializationJobStatus:
return self._status
def error(self) -> Optional[BaseException]:
return self._error
def should_be_retried(self) -> bool:
"""Ray jobs are generally not retried by default."""
return False
def url(self) -> Optional[str]:
"""Ray jobs don't have a specific URL."""
return None
def result(self) -> Optional[DAGValue]:
"""Get the result of the materialization job."""
return self._result