Skip to content

Commit 051dc7e

Browse files
Improved dataset merging logic.
Started working on improving the dataset and feature set merging logic. This commit improves dataset merging. It also includes updated tests for the changed code. Apart from that, it also includes some minor refactoring and code cleanup.
1 parent f03ecc0 commit 051dc7e

15 files changed

Lines changed: 353 additions & 117 deletions

File tree

docs/roadmap.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,7 @@ _Last updated: 2026-03-23_
2424
- Status: done (2026-03-23)
2525
- Importance: medium
2626
- [ ] Improve dataset and feature set merging logic
27-
- Status: planned
27+
- Status: in progress
2828
- Importance: medium
2929
- [ ] Create infer.py
3030
- Status: planned

ml/data/merge/merge_dataset_into_main.py

Lines changed: 83 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -3,19 +3,52 @@
33
import logging
44
from pathlib import Path
55

6+
import networkx as nx
67
import pandas as pd
78

89
from ml.data.validation.validate_data import validate_data
9-
from ml.exceptions import DataError
10+
from ml.exceptions import ConfigError, DataError
11+
from ml.feature_freezing.freeze_strategies.tabular.config.models import (
12+
DatasetConfig,
13+
MergeHow,
14+
MergeValidate,
15+
)
1016
from ml.utils.loaders import load_json
1117

1218
logger = logging.getLogger(__name__)
1319

20+
def build_dataset_dag(datasets: list[DatasetConfig]) -> list[str]:
21+
"""Build DAG and return topological merge order."""
22+
G = nx.DiGraph()
23+
for ds in datasets:
24+
G.add_node(ds.name)
25+
26+
for i, ds1 in enumerate(datasets):
27+
keys1 = set(ds1.merge_key if isinstance(ds1.merge_key, list) else [ds1.merge_key])
28+
for j, ds2 in enumerate(datasets):
29+
if i >= j:
30+
continue
31+
keys2 = set(ds2.merge_key if isinstance(ds2.merge_key, list) else [ds2.merge_key])
32+
if keys1 & keys2:
33+
G.add_edge(ds1.name, ds2.name)
34+
35+
try:
36+
merge_order = list(nx.topological_sort(G))
37+
except nx.NetworkXUnfeasible:
38+
msg = "Cyclic merge dependency detected among datasets."
39+
logger.error(msg)
40+
raise ConfigError(msg)
41+
42+
logger.debug(f"DAG-based merge order: {merge_order}")
43+
return merge_order
44+
1445
def merge_dataset_into_main(
1546
data: pd.DataFrame,
1647
df: pd.DataFrame,
1748
*,
18-
merge_key: str,
49+
merge_key: str | list[str],
50+
merge_how: MergeHow = "inner",
51+
merge_validate: MergeValidate = "m:m",
1952
dataset_name: str,
2053
dataset_version: str,
2154
dataset_snapshot_path: Path,
@@ -26,7 +59,9 @@ def merge_dataset_into_main(
2659
Args:
2760
data: Main dataframe accumulated from previous merges.
2861
df: New dataset dataframe to merge into ``data``.
29-
merge_key: Key column used for inner merge.
62+
merge_key: Key column(s) used for inner merge.
63+
merge_how: Type of merge to perform.
64+
merge_validate: Merge validation for row explosion protection.
3065
dataset_name: Logical dataset name used for logging/errors.
3166
dataset_version: Dataset version string.
3267
dataset_snapshot_path: Snapshot directory containing dataset metadata.
@@ -44,31 +79,58 @@ def merge_dataset_into_main(
4479
dataset-level merge diagnostics/warnings.
4580
"""
4681

47-
if merge_key not in df.columns:
48-
msg = f"Dataset {dataset_name} is missing merge key '{merge_key}' column."
82+
missing_keys_in_df = set(merge_key) - set(df.columns)
83+
if missing_keys_in_df:
84+
msg = f"Dataset {dataset_name} is missing merge key(s): {missing_keys_in_df}"
4985
logger.error(msg)
5086
raise DataError(msg)
51-
if merge_key not in data.columns and not data.empty:
52-
msg = f"Merge key '{merge_key}' column from dataset {dataset_name} not found in the main dataset for merging."
87+
88+
if not data.empty:
89+
missing_keys_in_main = set(merge_key) - set(data.columns)
90+
if missing_keys_in_main:
91+
msg = f"Merge key(s) {missing_keys_in_main} not found in main dataset when merging {dataset_name}"
92+
logger.error(msg)
93+
raise DataError(msg)
94+
95+
original_rows = len(data) if not data.empty else 0
96+
97+
overlapping_cols = set(data.columns) & set(df.columns) - set(merge_key)
98+
if overlapping_cols:
99+
logger.warning(f"Dropping overlapping columns {overlapping_cols} from dataset {dataset_name} before merge")
100+
df = df.drop(columns=overlapping_cols)
101+
102+
# Merge with row explosion protection
103+
try:
104+
if data.empty:
105+
data = df
106+
else:
107+
data = pd.merge(
108+
data,
109+
df,
110+
how=merge_how,
111+
on=merge_key,
112+
validate=merge_validate,
113+
suffixes=("", "_dup"),
114+
)
115+
except pd.errors.MergeError as e:
116+
msg = f"Merge failed for dataset {dataset_name} v{dataset_version}: {e}"
53117
logger.error(msg)
54118
raise DataError(msg)
55119

120+
if len(data) < original_rows:
121+
logger.warning(
122+
f"Row count decreased after merging {dataset_name}: {original_rows} -> {len(data)}"
123+
)
124+
125+
if merge_key:
126+
data = data.sort_values(by=merge_key).reset_index(drop=True)
127+
56128
dataset_metadata = load_json(dataset_snapshot_path / "metadata.json")
57129
data_hash = validate_data(data_path=dataset_path, metadata=dataset_metadata)
58130

59-
logger.debug(f"Starting to merge dataset {dataset_name} v{dataset_version} snapshot {dataset_snapshot_path.name} with shape {df.shape} into the main dataset with shape {data.shape}")
60-
61-
if data.empty:
62-
data = df
63-
else:
64-
overlapping_cols = set(data.columns) & set(df.columns) - {merge_key}
65-
if overlapping_cols:
66-
logger.warning(f"Overlapping columns found in dataset {dataset_name}: {overlapping_cols}. Dropping these columns from the new dataset before merge.")
67-
df = df.drop(columns=overlapping_cols)
68-
data = pd.merge(data, df, how="inner", on=merge_key)
69-
if data.empty:
70-
msg = f"Merged dataset is empty after merging with {dataset_name}. Please check the '{merge_key}' alignment across datasets."
71-
logger.error(msg)
72-
raise DataError(msg)
131+
logger.debug(
132+
f"Merged dataset {dataset_name} v{dataset_version} snapshot {dataset_snapshot_path.name} "
133+
f"into main dataframe, resulting shape: {data.shape}"
134+
)
73135

74136
return data, data_hash

ml/feature_freezing/freeze_strategies/tabular/config/models.py

Lines changed: 57 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -11,15 +11,70 @@
1111

1212
logger = logging.getLogger(__name__)
1313

14+
MergeHow = Literal["left", "right", "inner", "outer", "cross"]
15+
16+
MergeValidate = Literal[
17+
"one_to_one",
18+
"1:1",
19+
"one_to_many",
20+
"1:m",
21+
"many_to_one",
22+
"m:1",
23+
"many_to_many",
24+
"m:m",
25+
]
26+
1427
class DatasetConfig(BaseModel):
1528
"""Source dataset definition used for feature freezing ingestion."""
1629

1730
ref: str = Field("data/processed", description="Reference path for the dataset, e.g., 'data/processed'")
1831
name: str = Field(..., description="Name of the dataset, e.g., 'hotel_bookings'")
1932
version: str = Field(..., description="Version of the dataset, e.g., 'v1'")
2033
format: Literal["csv", "parquet"]
21-
merge_key: str = Field("row_id", description="Key to merge datasets on, default is 'row_id'")
22-
path_suffix: str = Field("data.{format}", description="Suffix for the dataset file, supports {format} placeholder")
34+
merge_key: str | list[str] = Field(
35+
"row_id", description="Key(s) to merge datasets on, default is 'row_id'"
36+
)
37+
merge_how: MergeHow = Field(
38+
"inner", description="Merge type, default 'inner'"
39+
)
40+
merge_validate: MergeValidate = Field(
41+
"m:m", description="Merge validation for row explosion protection"
42+
)
43+
path_suffix: str = Field(
44+
"data.{format}", description="Suffix for the dataset file, supports {format} placeholder"
45+
)
46+
47+
@field_validator("merge_key", mode="before")
48+
def ensure_merge_key_list(cls, v):
49+
"""Always convert merge_key to a list internally."""
50+
if isinstance(v, str):
51+
return [v]
52+
if isinstance(v, list) and all(isinstance(i, str) for i in v):
53+
return v
54+
raise ConfigError("merge_key must be a string or a list of strings")
55+
56+
@field_validator("merge_how")
57+
def validate_merge_how(cls, v):
58+
if v not in {"inner", "left", "right", "outer"}:
59+
raise ConfigError(f"Invalid merge_how: {v}")
60+
return v
61+
62+
@field_validator("merge_validate")
63+
def normalize_merge_validate(cls, v):
64+
"""Normalize merge_validate to pandas-compatible strings."""
65+
mapping = {
66+
"one_to_one": "1:1",
67+
"1:1": "1:1",
68+
"one_to_many": "1:m",
69+
"1:m": "1:m",
70+
"many_to_one": "m:1",
71+
"m:1": "m:1",
72+
"many_to_many": "m:m",
73+
"m:m": "m:m",
74+
}
75+
if v not in mapping:
76+
raise ConfigError(f"Invalid merge_validate: {v}")
77+
return mapping[v]
2378

2479
class FeatureRolesConfig(BaseModel):
2580
"""Feature role partitioning for validation and downstream usage."""

ml/feature_freezing/utils/data_loader.py

Lines changed: 43 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55

66
import pandas as pd
77

8-
from ml.data.merge.merge_dataset_into_main import merge_dataset_into_main
8+
from ml.data.merge.merge_dataset_into_main import build_dataset_dag, merge_dataset_into_main
99
from ml.exceptions import ConfigError, DataError
1010
from ml.feature_freezing.freeze_strategies.tabular.config.models import TabularFeaturesConfig
1111
from ml.snapshot_bindings.extraction.get_snapshot_binding import get_and_validate_snapshot_binding
@@ -16,10 +16,12 @@
1616

1717
logger = logging.getLogger(__name__)
1818

19+
20+
1921
def load_data_with_lineage(
20-
config: TabularFeaturesConfig,
21-
snapshot_binding_key: str | None
22-
) -> tuple[pd.DataFrame, list[DataLineageEntry]]:
22+
config: TabularFeaturesConfig,
23+
snapshot_binding_key: str | None = None
24+
) -> tuple[pd.DataFrame, list[DataLineageEntry]]:
2325
"""Load configured datasets, merge them, and build lineage entries.
2426
2527
Args:
@@ -30,9 +32,6 @@ def load_data_with_lineage(
3032
dataset lineage metadata.
3133
"""
3234

33-
data = pd.DataFrame()
34-
data_lineage = []
35-
3635
if not config.data:
3736
msg = "No datasets specified in the configuration."
3837
logger.error(msg)
@@ -41,70 +40,71 @@ def load_data_with_lineage(
4140
datasets_binding = None
4241
if snapshot_binding_key:
4342
snapshot_binding_config = get_and_validate_snapshot_binding(
44-
snapshot_binding_key,
45-
expect_dataset_bindings=True
43+
snapshot_binding_key, expect_dataset_bindings=True
4644
)
4745
datasets_binding = snapshot_binding_config.datasets
4846

49-
for dataset in config.data:
50-
dataset_versions = datasets_binding.get(dataset.name) if datasets_binding else None
51-
dataset_snapshot_binding = dataset_versions.get(dataset.version) if dataset_versions else None
47+
dataset_dict = {ds.name: ds for ds in config.data}
48+
merge_order = build_dataset_dag(config.data)
49+
50+
merged_data = pd.DataFrame()
51+
data_lineage: list[DataLineageEntry] = []
52+
53+
for ds_name in merge_order:
54+
ds = dataset_dict[ds_name]
55+
dataset_versions = datasets_binding.get(ds.name) if datasets_binding else None
56+
dataset_snapshot_binding = dataset_versions.get(ds.version) if dataset_versions else None
57+
5258
if snapshot_binding_key:
5359
if not dataset_snapshot_binding:
54-
msg = f"No snapshot binding found for dataset {dataset.name} version {dataset.version} under snapshot binding key {snapshot_binding_key}."
60+
msg = f"No snapshot binding found for {ds.name} {ds.version} under {snapshot_binding_key}"
5561
logger.error(msg)
5662
raise ConfigError(msg)
5763
dataset_snapshot = dataset_snapshot_binding.snapshot
58-
dataset_snapshot_path = Path(dataset.ref) / dataset.name / dataset.version / dataset_snapshot
64+
dataset_snapshot_path = Path(ds.ref) / ds.name / ds.version / dataset_snapshot
5965
else:
60-
dataset_path = Path(dataset.ref) / dataset.name / dataset.version
61-
dataset_snapshot_path = get_latest_snapshot_path(dataset_path)
66+
dataset_path_base = Path(ds.ref) / ds.name / ds.version
67+
dataset_snapshot_path = get_latest_snapshot_path(dataset_path_base)
6268

63-
if dataset.format not in HASH_LOADER_REGISTRY:
64-
msg = f"Unsupported data format for loading and hashing: {dataset.format}"
65-
logger.error(msg)
66-
raise ConfigError(msg)
67-
68-
dataset_path = dataset_snapshot_path / dataset.path_suffix.format(format=dataset.format)
69+
dataset_path = dataset_snapshot_path / ds.path_suffix.format(format=ds.format)
6970
if not dataset_path.exists():
7071
msg = f"Dataset file not found at expected path: {dataset_path}"
7172
logger.error(msg)
7273
raise DataError(msg)
73-
df = read_data(dataset.format, dataset_path)
74-
merge_key = dataset.merge_key
74+
df = read_data(ds.format, dataset_path)
7575

76-
data, data_hash = merge_dataset_into_main(
77-
data=data,
76+
merged_data, data_hash = merge_dataset_into_main(
77+
data=merged_data,
7878
df=df,
79-
merge_key=merge_key,
80-
dataset_name=dataset.name,
81-
dataset_version=dataset.version,
79+
merge_key=ds.merge_key,
80+
merge_how=ds.merge_how,
81+
merge_validate=ds.merge_validate,
82+
dataset_name=ds.name,
83+
dataset_version=ds.version,
8284
dataset_snapshot_path=dataset_snapshot_path,
8385
dataset_path=dataset_path,
8486
)
8587

86-
loader_validation_hash = HASH_LOADER_REGISTRY[dataset.format](dataset_path)
88+
loader_validation_hash = HASH_LOADER_REGISTRY[ds.format](dataset_path)
8789

8890
entry = DataLineageEntry(
89-
ref=dataset.ref,
90-
name=dataset.name,
91-
version=dataset.version,
92-
format=dataset.format,
93-
path_suffix=dataset.path_suffix,
94-
merge_key=dataset.merge_key,
91+
ref=ds.ref,
92+
name=ds.name,
93+
version=ds.version,
94+
format=ds.format,
95+
path_suffix=ds.path_suffix,
96+
merge_key=ds.merge_key,
97+
merge_how=ds.merge_how,
98+
merge_validate=ds.merge_validate,
9599
snapshot_id=dataset_snapshot_path.name,
96100
path=dataset_path.as_posix(),
97101
loader_validation_hash=loader_validation_hash,
98102
data_hash=data_hash,
99103
row_count=len(df),
100104
column_count=len(df.columns),
101105
)
102-
103106
data_lineage.append(entry)
107+
logger.debug(f"Loaded dataset {ds.name} {ds.version} snapshot {dataset_snapshot_path.name} loader hash {loader_validation_hash}")
104108

105-
logger.debug(f"Loaded dataset {dataset.name} {dataset.version} snapshot {dataset_snapshot_path.name} "
106-
f"file {dataset_path} loader hash {loader_validation_hash}")
107-
108-
logger.info(f"Completed loading {len(config.data)} dataframes. Final merged dataframe shape: {data.shape}. Lineage: {data_lineage}")
109-
110-
return data, data_lineage
109+
logger.info(f"Completed loading {len(config.data)} datasets. Final merged dataframe shape: {merged_data.shape}.")
110+
return merged_data, data_lineage

0 commit comments

Comments
 (0)