Skip to content

Commit ffab2a2

Browse files
fix(learning): text dtypes compatible in replay; merge list memory patterns (#419)
* fix(learning): treat text dtypes as compatible in profile replay The drift gate only treated object as text, so a profile learned on a categorical or string column and replayed on the same values as object (or the reverse) reported mild drift and dropped that column's value maps and hints. A text categorical replayed on numeric data was also not flagged. Compare text-ness with _util.is_text_dtype on the frame's dtype and a name-based equivalent for the stored learned dtype. object, string, Arrow string and text categoricals are interchangeable; text vs non-text is still reported as drift. * fix(learning): merge list and scalar memory value patterns union_min_precision and error_on_conflict assumed every entry of the embedded memory's value_patterns was a {raw: clean} mapping and raised ValueError when memory held the semantic_repairs list. Union list patterns with de-duplication. A semantic repair both sides propose differently is dropped and recorded under union_min_precision and raises ProfileMergeError under error_on_conflict; differing scalar patterns are handled the same way. prefer_self/prefer_other still copy the preferred side's memory. * docs(learning): note that masked tokens are per profile The masking salt comes from the frame signature (row count, dtypes and a head sample), so the same values can mask to different tokens in different profiles and merged profiles cannot match each other's masked entries.
1 parent 13ff83e commit ffab2a2

5 files changed

Lines changed: 354 additions & 16 deletions

File tree

‎src/freshdata/learning/merge.py‎

Lines changed: 108 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
from __future__ import annotations
1616

1717
import copy
18+
import json
1819
from collections.abc import Mapping
1920
from dataclasses import dataclass, field
2021
from datetime import datetime, timezone
@@ -247,6 +248,11 @@ def merge_profiles(self_profile: Any, other_profile: Any, *, strategy: Strategy)
247248
)
248249
conflicts.extend(map_conflicts_found)
249250

251+
memory, memory_conflicts = _merge_memory_checked(
252+
self_profile.memory, other_profile.memory, strategy
253+
)
254+
conflicts.extend(memory_conflicts)
255+
250256
if strategy == "error_on_conflict" and conflicts:
251257
raise ProfileMergeError(
252258
"profiles conflict; refusing to merge:\n - " + "\n - ".join(conflicts)
@@ -255,7 +261,6 @@ def merge_profiles(self_profile: Any, other_profile: Any, *, strategy: Strategy)
255261
notes.extend(f"conflict ({_resolution_word(strategy)}): {c}" for c in conflicts)
256262

257263
examples = _merge_examples(self_profile.examples, other_profile.examples)
258-
memory = _merge_memory(self_profile.memory, other_profile.memory, strategy)
259264

260265
privacy_mode = _merged_privacy(self_profile, other_profile, strategy)
261266
base = other_profile if strategy == "prefer_other" else self_profile
@@ -469,21 +474,110 @@ def _merge_memory(a: Any, b: Any, strategy: str) -> Any:
469474
or any later change to the merged profile, rewrite the parents' replay
470475
behaviour and saved bytes.
471476
"""
477+
return _merge_memory_checked(a, b, strategy)[0]
478+
479+
480+
def _merge_memory_checked(a: Any, b: Any, strategy: str) -> tuple[Any, list[str]]:
481+
""":func:`_merge_memory` plus the conflicts found in non-mapping patterns.
482+
483+
``prefer_self``/``prefer_other`` copy the preferred side's memory whole, so
484+
they never conflict. The union keeps self's memory and folds in the other
485+
side's patterns: per-column ``{raw: clean}`` mappings drop a disagreeing
486+
raw value (the value-map merge already reports those), while list patterns
487+
such as ``semantic_repairs`` are unioned and de-duplicated, and a repair
488+
both sides propose differently is dropped and reported as a conflict.
489+
"""
472490
if a is None:
473-
return copy.deepcopy(b)
491+
return copy.deepcopy(b), []
474492
if b is None or strategy == "prefer_self":
475-
return copy.deepcopy(a)
493+
return copy.deepcopy(a), []
476494
if strategy == "prefer_other":
477-
return copy.deepcopy(b)
478-
# union: keep self's memory but fold in non-conflicting value patterns.
495+
return copy.deepcopy(b), []
479496
merged = copy.deepcopy(a)
480-
merged_patterns = {c: dict(p) for c, p in merged.value_patterns.items()}
481-
for column, patterns in b.value_patterns.items():
482-
target = merged_patterns.setdefault(column, {})
483-
for raw, clean in patterns.items():
484-
if raw in target and str(target[raw]) != str(clean):
485-
del target[raw] # conflicting mapping: drop from replay
486-
elif raw not in target:
487-
target[raw] = copy.deepcopy(clean)
497+
conflicts: list[str] = []
498+
merged_patterns: dict[str, Any] = dict(merged.value_patterns)
499+
for key, theirs in b.value_patterns.items():
500+
if key not in merged_patterns:
501+
merged_patterns[key] = copy.deepcopy(theirs)
502+
continue
503+
ours = merged_patterns[key]
504+
if isinstance(ours, Mapping) and isinstance(theirs, Mapping):
505+
target = dict(ours)
506+
for raw, clean in theirs.items():
507+
if raw in target and str(target[raw]) != str(clean):
508+
del target[raw] # conflicting mapping: drop from replay
509+
elif raw not in target:
510+
target[raw] = copy.deepcopy(clean)
511+
merged_patterns[key] = target
512+
elif isinstance(ours, list) and isinstance(theirs, list):
513+
merged_patterns[key], found = _merge_pattern_lists(str(key), ours, theirs)
514+
conflicts.extend(found)
515+
elif _canonical(ours) != _canonical(theirs):
516+
del merged_patterns[key]
517+
conflicts.append(
518+
f"memory value_patterns '{key}': {type(ours).__name__} vs "
519+
f"{type(theirs).__name__} values differ"
520+
)
488521
merged.value_patterns = merged_patterns
489-
return merged
522+
return merged, conflicts
523+
524+
525+
_REPAIR_IDENTITY_FIELDS = ("column", "issue_type", "expert", "raw_value")
526+
527+
528+
def _canonical(value: object) -> str:
529+
try:
530+
return json.dumps(value, sort_keys=True, default=repr)
531+
except (TypeError, ValueError): # mixed-type keys cannot be sorted
532+
return repr(value)
533+
534+
535+
def _repair_identity(entry: object) -> tuple[str, str] | None:
536+
"""``(identity, proposal)`` for a repair record, or None for any other entry."""
537+
if not isinstance(entry, Mapping) or "proposed_value" not in entry:
538+
return None
539+
identity = _canonical({f: entry.get(f) for f in _REPAIR_IDENTITY_FIELDS})
540+
proposal = _canonical([entry.get("proposed_value"), entry.get("proposed_type", "str")])
541+
return identity, proposal
542+
543+
544+
def _merge_pattern_lists(
545+
key: str, ours: list[Any], theirs: list[Any]
546+
) -> tuple[list[Any], list[str]]:
547+
"""Union of two list patterns (e.g. ``semantic_repairs``), first copy wins.
548+
549+
Repair records are the same entry when column, issue type, expert and raw
550+
value match. When the two sides propose different values for one entry,
551+
it is dropped from both and reported; other entries de-duplicate by value.
552+
"""
553+
proposals: tuple[dict[str, set[str]], dict[str, set[str]]] = ({}, {})
554+
labels: dict[str, str] = {}
555+
for side, entries in zip(proposals, (ours, theirs)):
556+
for entry in entries:
557+
found = _repair_identity(entry)
558+
if found is None:
559+
continue
560+
side.setdefault(found[0], set()).add(found[1])
561+
labels.setdefault(
562+
found[0], f"column {entry.get('column')!r}, issue {entry.get('issue_type')!r}"
563+
)
564+
conflicting = {
565+
identity
566+
for identity in proposals[0].keys() & proposals[1].keys()
567+
if proposals[0][identity] != proposals[1][identity]
568+
}
569+
conflicts = [
570+
f"memory value_patterns '{key}': {labels[identity]} proposes different values"
571+
for identity in sorted(conflicting, key=lambda i: labels[i])
572+
]
573+
merged: list[Any] = []
574+
seen: set[str] = set()
575+
for entry in (*ours, *theirs):
576+
found = _repair_identity(entry)
577+
if found is not None and found[0] in conflicting:
578+
continue
579+
marker = _canonical(found) if found is not None else _canonical(entry)
580+
if marker not in seen:
581+
seen.add(marker)
582+
merged.append(copy.deepcopy(entry))
583+
return merged, conflicts

‎src/freshdata/learning/privacy.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,14 @@
1414
Masking is deterministic per profile: the salt is derived from the training
1515
pair's dataset signature so re-learning the same pair yields identical
1616
tokens (and therefore identical profile hashes).
17+
18+
Tokens are per profile, not global. The dataset signature covers the messy
19+
frame's row count, column names and dtypes, and a sample of its first rows,
20+
so the same values learned from a frame that differs in any of those (for
21+
example a ``category`` column instead of ``object``, or more rows) get
22+
different tokens. Masked tokens are therefore only comparable within one
23+
profile: a merged profile cannot match one parent's masked entries against
24+
the other's, so they are neither combined nor reported as conflicts.
1725
"""
1826

1927
from __future__ import annotations

‎src/freshdata/learning/replay.py‎

Lines changed: 29 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@
2727

2828
import pandas as pd
2929

30+
from .._util import is_text_dtype
3031
from .profile import LearningProfile, load_profile
3132

3233
__all__ = [
@@ -100,6 +101,31 @@ def _profile_schema(profile: LearningProfile) -> dict[str, str]:
100101
return {}
101102

102103

104+
#: Arrow text types, as spelled before the ``[pyarrow]`` suffix of a dtype name.
105+
_ARROW_TEXT_TYPES = ("string", "large_string", "string_view")
106+
107+
108+
def _is_text_dtype_name(name: str) -> bool:
109+
""":func:`is_text_dtype` for a stored dtype *name* (profiles keep ``str(dtype)``).
110+
111+
A learned ``"category"`` carries no categories, so it counts as text, as
112+
it does in ``CleaningMemory`` signatures; a frame's categorical is judged
113+
by its real categories.
114+
"""
115+
if name in ("category", "str"): # "str": the pandas 3 default string dtype
116+
return True
117+
try:
118+
return is_text_dtype(pd.api.types.pandas_dtype(name))
119+
except (TypeError, ValueError, ImportError, NotImplementedError):
120+
pass
121+
base, _, backend = name.partition("[")
122+
if backend != "pyarrow]":
123+
return False
124+
if base.startswith("dictionary<values="):
125+
base = base[len("dictionary<values=") :].split(",", 1)[0]
126+
return base in _ARROW_TEXT_TYPES
127+
128+
103129
def _referenced_columns(profile: LearningProfile) -> set[str]:
104130
columns = {r.column for r in profile.rules if r.column}
105131
columns.update(profile.value_maps)
@@ -136,8 +162,9 @@ def check_profile_drift(df: pd.DataFrame, profile: LearningProfile) -> ProfileRe
136162
learned = schema.get(column)
137163
if learned is None:
138164
continue
139-
actual = str(df[column].dtype)
140-
if learned != actual and (learned == "object") != (actual == "object"):
165+
dtype = df[column].dtype
166+
actual = str(dtype)
167+
if learned != actual and _is_text_dtype_name(learned) != is_text_dtype(dtype):
141168
dtype_incompatible.append(f"{column} ({learned} -> {actual})")
142169
if dtype_incompatible:
143170
reasons.append("dtype changed for: " + ", ".join(dtype_incompatible[:5]))

‎tests/learning/test_merge_diff.py‎

Lines changed: 108 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,8 @@
22

33
from __future__ import annotations
44

5+
import copy
6+
57
import pytest
68

79
from freshdata.learning import learn, load_profile, save_profile
@@ -93,6 +95,112 @@ def test_unknown_strategy_rejected(self, base_profile):
9395
merge_profiles(base_profile, base_profile, strategy="bogus")
9496

9597

98+
_STRATEGIES = ("union_min_precision", "prefer_self", "prefer_other", "error_on_conflict")
99+
100+
101+
def _repair(raw_value, proposed_value, **extra):
102+
return {
103+
"column": "status",
104+
"issue_type": "case",
105+
"expert": "case_normalization",
106+
"raw_value": raw_value,
107+
"proposed_value": proposed_value,
108+
"proposed_type": "str",
109+
**extra,
110+
}
111+
112+
113+
def _with_patterns(profile, **patterns):
114+
"""Deep copy of *profile* whose embedded memory carries extra value patterns."""
115+
copied = copy.deepcopy(profile)
116+
assert copied.memory is not None
117+
copied.memory.value_patterns.update(copy.deepcopy(patterns))
118+
return copied
119+
120+
121+
def _notes(profile):
122+
return list(profile.audit_info.notes)
123+
124+
125+
class TestMergeMemoryListPatterns:
126+
"""``semantic_repairs`` and other non-mapping memory patterns merge without crashing."""
127+
128+
@pytest.mark.parametrize("strategy", _STRATEGIES)
129+
def test_semantic_repairs_union_deduplicated(self, base_profile, strategy):
130+
shared, ours, theirs = _repair("SHIPPED", "shipped"), _repair("A", "a"), _repair("B", "b")
131+
a = _with_patterns(base_profile, semantic_repairs=[shared, ours])
132+
b = _with_patterns(base_profile, semantic_repairs=[dict(shared, confidence=0.5), theirs])
133+
134+
merged = merge_profiles(a, b, strategy=strategy)
135+
136+
repairs = merged.memory.value_patterns["semantic_repairs"]
137+
expected = {
138+
"prefer_self": [shared, ours],
139+
"prefer_other": [dict(shared, confidence=0.5), theirs],
140+
}.get(strategy, [shared, ours, theirs])
141+
assert repairs == expected
142+
assert not any("semantic_repairs" in n for n in _notes(merged))
143+
# The parents are untouched.
144+
assert a.memory.value_patterns["semantic_repairs"] == [shared, ours]
145+
assert repairs is not a.memory.value_patterns["semantic_repairs"]
146+
147+
@pytest.mark.parametrize("strategy", _STRATEGIES)
148+
def test_semantic_repairs_conflict(self, base_profile, strategy):
149+
agreed = _repair("SHIPPED", "shipped")
150+
ours, theirs = _repair("Deliverd", "delivered"), _repair("Deliverd", "returned")
151+
a = _with_patterns(base_profile, semantic_repairs=[ours, agreed])
152+
b = _with_patterns(base_profile, semantic_repairs=[theirs, agreed])
153+
154+
if strategy == "error_on_conflict":
155+
with pytest.raises(ProfileMergeError, match="semantic_repairs"):
156+
merge_profiles(a, b, strategy=strategy)
157+
return
158+
merged = merge_profiles(a, b, strategy=strategy)
159+
proposals = {
160+
r["raw_value"]: r["proposed_value"]
161+
for r in merged.memory.value_patterns["semantic_repairs"]
162+
}
163+
if strategy == "union_min_precision":
164+
assert proposals == {"SHIPPED": "shipped"}
165+
assert any("semantic_repairs" in n for n in _notes(merged))
166+
else:
167+
winner = "delivered" if strategy == "prefer_self" else "returned"
168+
assert proposals == {"Deliverd": winner, "SHIPPED": "shipped"}
169+
170+
@pytest.mark.parametrize("strategy", _STRATEGIES)
171+
def test_semantic_repairs_from_saved_memory(self, base_profile, strategy, tmp_path):
172+
a = _with_patterns(base_profile, semantic_repairs=[_repair("SHIPPED", "shipped")])
173+
b = _with_patterns(base_profile, semantic_repairs=[_repair("Deliverd", "delivered")])
174+
path_a, path_b = tmp_path / "a.fdprofile", tmp_path / "b.fdprofile"
175+
save_profile(a, path_a)
176+
save_profile(b, path_b)
177+
178+
merged = merge_profiles(load_profile(path_a), load_profile(path_b), strategy=strategy)
179+
180+
raw_values = [r["raw_value"] for r in merged.memory.value_patterns["semantic_repairs"]]
181+
expected = {"prefer_self": ["SHIPPED"], "prefer_other": ["Deliverd"]}
182+
assert raw_values == expected.get(strategy, ["SHIPPED", "Deliverd"])
183+
184+
@pytest.mark.parametrize("strategy", _STRATEGIES)
185+
def test_other_non_mapping_patterns(self, base_profile, strategy):
186+
a = _with_patterns(base_profile, tags=["x", "y"], version=1, note="same")
187+
b = _with_patterns(base_profile, tags=["y", "z"], version=2, note="same")
188+
189+
if strategy == "error_on_conflict":
190+
with pytest.raises(ProfileMergeError, match="version"):
191+
merge_profiles(a, b, strategy=strategy)
192+
return
193+
patterns = merge_profiles(a, b, strategy=strategy).memory.value_patterns
194+
if strategy == "union_min_precision":
195+
assert patterns["tags"] == ["x", "y", "z"]
196+
assert "version" not in patterns
197+
assert patterns["note"] == "same"
198+
else:
199+
side = a if strategy == "prefer_self" else b
200+
assert patterns["tags"] == side.memory.value_patterns["tags"]
201+
assert patterns["version"] == side.memory.value_patterns["version"]
202+
203+
96204
class TestMergedProfileIntegrity:
97205
def test_merged_id_differs_and_roundtrips(self, base_profile, conflicting_profile, tmp_path):
98206
merged = merge_profiles(base_profile, conflicting_profile, strategy="union_min_precision")

0 commit comments

Comments
 (0)