Note
Go to the end to download the full example code.
Custom Entity Metric#
Learn how to implement a custom entity-level distance metric by subclassing
EntityMetric.
An entity metric computes a scalar distance between two individual entities (atomic observations in a sequence). It is the basic building block used by most sequence-level metrics.
Minimal contract
Declare a
SETTINGS_CLASSdataclass (usetanat_utils.settings_dataclass).Implement
validate_entity(ent_a, ent_b): raiseTypeError/KeyErrorwhen the entities are incompatible with the metric.Implement
_compute(ent_a, ent_b): return a non-negativefloat.
The public __call__ in the base class invokes validate_entity then
_compute automatically.
Setup#
import polars as pl
from tanat_utils import settings_dataclass as dataclass
from tanat import build_events
from tanat.dataset import simulate_events
from tanat.metric.entity.base import EntityMetric
Data#
A small event pool with a score numeric feature and a status
categorical feature.
raw = simulate_events(n_ids=20, features=["score", "status"], seed=0)
pool = build_events(temporal_data=raw, id_column="id", time_column="time")
pool.cast_features({"score": pl.Float32})
┌─ Event SequenceStore
│
│ Step 1/4: Sorting & preparing data
│
│ Step 2/4: Building sequence index
│
│ Step 3/4: Writing entity & time index features
│
│ Step 4/4: Computing & writing metadata
│
└─ Done (20 sequences · 138 entities · 0.00s)
print(pool)
┌────────────────────────────────────────────────┐
│ EventSequencePool Summary │
└────────────────────────────────────────────────┘
Overview
─────────────────────────
Sequences 20
Store /home/runner/.tanat/_quick_event_6c691121
id_column id
Time Index
─────────────────────────
Type Datetime(time_unit='us', time_zone=None) [2000-01-03 17:54:05.937674 → 2024-11-15 14:02:43.302235]
Columns ['time']
t0 position=0, anchor=None
Entity Features (2)
─────────────────────────
• score Numerical [1.0 → 100.0]
• status String [len 1 → 1]
Define a custom entity metric#
We implement a normalised absolute difference on a numeric feature.
dist(a, b) = |a.score - b.score| / scale
@dataclass
class ScoreDiffSettings:
"""Settings for :class:`ScoreDiffEntityMetric`.
Args:
entity_feature: Numeric feature to compare.
scale: Normalisation constant (``1.0`` means raw absolute diff).
"""
entity_feature: str = "score"
scale: float = 1.0
class ScoreDiffEntityMetric(EntityMetric, register_name="score_diff"):
"""Normalised absolute difference on a numeric entity feature."""
SETTINGS_CLASS = ScoreDiffSettings
def __init__(self, entity_feature: str = "score", scale: float = 1.0) -> None:
super().__init__(
settings=ScoreDiffSettings(entity_feature=entity_feature, scale=scale)
)
# ------------------------------------------------------------------
# Required implementations
# ------------------------------------------------------------------
def validate_entity(self, ent_a, ent_b=None) -> None:
"""Check that both entities expose the expected numeric feature."""
self._validate_entity_instance(ent_a, ent_b)
feat = self.settings.entity_feature
for ent in (e for e in (ent_a, ent_b) if e is not None):
if feat not in ent.data():
raise KeyError(
f"Feature {feat!r} not found in entity. "
f"Available: {list(ent.data().keys())}"
)
def _compute(self, ent_a, ent_b) -> float:
"""Return normalised absolute difference of the configured feature."""
feat = self.settings.entity_feature
val_a = float(ent_a[feat])
val_b = float(ent_b[feat])
return abs(val_a - val_b) / self.settings.scale
Instantiate and inspect#
metric = ScoreDiffEntityMetric(entity_feature="score", scale=100.0)
print(metric)
ScoreDiffEntityMetric(settings=ScoreDiffSettings(entity_feature='score', scale=100.0))
Compute a distance between two entities#
ids = pool.unique_ids
ent_a = pool[ids[0]][0]
ent_b = pool[ids[1]][0]
print(f"score A : {ent_a['score']}")
print(f"score B : {ent_b['score']}")
print(f"distance: {metric(ent_a, ent_b):.4f}")
score A : 28.0
score B : 73.0
distance: 0.4500
Compute distances over several pairs#
print("Sample pairwise distances")
print("-" * 40)
for i in range(5):
ea = pool[ids[i]][0]
eb = pool[ids[i + 1]][0]
d = metric(ea, eb)
print(f" {ea['score']:6.1f} vs {eb['score']:6.1f} -> {d:.4f}")
Sample pairwise distances
----------------------------------------
28.0 vs 73.0 -> 0.4500
73.0 vs 30.0 -> 0.4300
30.0 vs 1.0 -> 0.2900
1.0 vs 62.0 -> 0.6100
62.0 vs 81.0 -> 0.1900
Use the custom metric inside a sequence metric#
Any EntityMetric can be passed as the
entity_metric argument to sequence-level metrics such as
LinearPairwiseSequenceMetric.
from tanat.metric.sequence import LinearPairwiseSequenceMetric
lp = LinearPairwiseSequenceMetric(entity_metric=metric)
print(lp)
LinearPairwiseSequenceMetric(settings=LinearPairwiseSettings(entity_metric=ScoreDiffEntityMetric(settings=ScoreDiffSettings(entity_feature='score', scale=100.0)), agg_fun='mean', padding_penalty=None))
seq_a = pool[ids[0]]
seq_b = pool[ids[1]]
dist = lp(seq_a, seq_b)
print(f"LinearPairwise distance: {dist:.4f}")
LinearPairwise distance: 0.3025
dm = lp.compute_matrix(pool)
dm.to_frame().head()
# .. note::
#
# Such an custom metric is not optimized for large datasets. TanaT has mechanisms
# based on the Numba framework.
# For implementing an
# computationaly efficient metric, we invite the developers to contact the core
# developer team or to look at the code of the implemented entity metrics (especially
# hamming).
┌─ LinearPairwiseSequenceMetric
│
│ Pairs: 0%| | 0/400 [00:00<?, ?it/s]
│ Pairs: 1%| | 3/400 [00:00<00:16, 23.77it/s]
│ Pairs: 2%|▏ | 8/400 [00:00<00:10, 36.83it/s]
│ Pairs: 3%|▎ | 12/400 [00:00<00:12, 30.64it/s]
│ Pairs: 4%|▍ | 16/400 [00:00<00:13, 27.92it/s]
│ Pairs: 5%|▍ | 19/400 [00:00<00:13, 27.29it/s]
│ Pairs: 6%|▌ | 22/400 [00:00<00:14, 25.86it/s]
│ Pairs: 6%|▋ | 26/400 [00:00<00:12, 29.02it/s]
│ Pairs: 8%|▊ | 31/400 [00:01<00:11, 31.30it/s]
│ Pairs: 9%|▉ | 35/400 [00:01<00:12, 29.20it/s]
│ Pairs: 10%|▉ | 38/400 [00:01<00:12, 27.96it/s]
│ Pairs: 10%|█ | 41/400 [00:01<00:12, 27.63it/s]
│ Pairs: 11%|█▏ | 45/400 [00:01<00:12, 29.02it/s]
│ Pairs: 12%|█▎ | 50/400 [00:01<00:10, 33.77it/s]
│ Pairs: 14%|█▎ | 54/400 [00:01<00:10, 31.70it/s]
│ Pairs: 14%|█▍ | 58/400 [00:01<00:11, 30.31it/s]
│ Pairs: 16%|█▌ | 62/400 [00:02<00:10, 30.88it/s]
│ Pairs: 17%|█▋ | 67/400 [00:02<00:09, 34.70it/s]
│ Pairs: 18%|█▊ | 72/400 [00:02<00:08, 37.04it/s]
│ Pairs: 19%|█▉ | 76/400 [00:02<00:08, 37.14it/s]
│ Pairs: 20%|██ | 80/400 [00:02<00:08, 36.99it/s]
│ Pairs: 21%|██ | 84/400 [00:02<00:08, 36.99it/s]
│ Pairs: 22%|██▏ | 89/400 [00:02<00:07, 40.54it/s]
│ Pairs: 24%|██▎ | 94/400 [00:02<00:07, 39.66it/s]
│ Pairs: 25%|██▍ | 99/400 [00:02<00:07, 39.14it/s]
│ Pairs: 26%|██▋ | 105/400 [00:03<00:06, 43.46it/s]
│ Pairs: 28%|██▊ | 112/400 [00:03<00:05, 48.68it/s]
│ Pairs: 30%|██▉ | 119/400 [00:03<00:05, 52.95it/s]
│ Pairs: 32%|███▏ | 126/400 [00:03<00:04, 56.00it/s]
│ Pairs: 33%|███▎ | 132/400 [00:03<00:04, 56.89it/s]
│ Pairs: 35%|███▍ | 139/400 [00:03<00:04, 58.00it/s]
│ Pairs: 36%|███▋ | 145/400 [00:03<00:04, 58.21it/s]
│ Pairs: 38%|███▊ | 152/400 [00:03<00:04, 58.88it/s]
│ Pairs: 40%|███▉ | 158/400 [00:03<00:04, 59.01it/s]
│ Pairs: 41%|████ | 164/400 [00:04<00:04, 56.32it/s]
│ Pairs: 42%|████▎ | 170/400 [00:04<00:04, 55.36it/s]
│ Pairs: 44%|████▍ | 176/400 [00:04<00:04, 53.01it/s]
│ Pairs: 46%|████▌ | 182/400 [00:04<00:04, 45.86it/s]
│ Pairs: 47%|████▋ | 187/400 [00:04<00:04, 44.41it/s]
│ Pairs: 48%|████▊ | 192/400 [00:04<00:05, 38.68it/s]
│ Pairs: 49%|████▉ | 197/400 [00:05<00:06, 33.26it/s]
│ Pairs: 50%|█████ | 201/400 [00:05<00:06, 30.65it/s]
│ Pairs: 51%|█████▏ | 205/400 [00:05<00:06, 30.67it/s]
│ Pairs: 52%|█████▎ | 210/400 [00:05<00:05, 34.25it/s]
│ Pairs: 54%|█████▎ | 214/400 [00:05<00:05, 31.72it/s]
│ Pairs: 55%|█████▍ | 218/400 [00:05<00:06, 29.63it/s]
│ Pairs: 56%|█████▌ | 222/400 [00:05<00:06, 27.93it/s]
│ Pairs: 56%|█████▋ | 226/400 [00:05<00:05, 30.34it/s]
│ Pairs: 58%|█████▊ | 231/400 [00:06<00:05, 32.07it/s]
│ Pairs: 59%|█████▉ | 235/400 [00:06<00:05, 29.17it/s]
│ Pairs: 60%|█████▉ | 239/400 [00:06<00:05, 28.31it/s]
│ Pairs: 60%|██████ | 242/400 [00:06<00:05, 27.30it/s]
│ Pairs: 62%|██████▏ | 246/400 [00:06<00:05, 29.85it/s]
│ Pairs: 63%|██████▎ | 251/400 [00:06<00:04, 32.95it/s]
│ Pairs: 64%|██████▍ | 255/400 [00:06<00:04, 31.34it/s]
│ Pairs: 65%|██████▍ | 259/400 [00:07<00:04, 30.22it/s]
│ Pairs: 66%|██████▌ | 263/400 [00:07<00:04, 29.44it/s]
│ Pairs: 67%|██████▋ | 268/400 [00:07<00:03, 34.01it/s]
│ Pairs: 68%|██████▊ | 272/400 [00:07<00:03, 32.85it/s]
│ Pairs: 69%|██████▉ | 276/400 [00:07<00:03, 31.21it/s]
│ Pairs: 70%|███████ | 280/400 [00:07<00:03, 30.26it/s]
│ Pairs: 71%|███████ | 284/400 [00:07<00:03, 29.10it/s]
│ Pairs: 72%|███████▎ | 290/400 [00:08<00:03, 32.87it/s]
│ Pairs: 74%|███████▎ | 294/400 [00:08<00:03, 29.80it/s]
│ Pairs: 74%|███████▍ | 298/400 [00:08<00:03, 27.78it/s]
│ Pairs: 75%|███████▌ | 301/400 [00:08<00:03, 26.73it/s]
│ Pairs: 76%|███████▌ | 304/400 [00:08<00:03, 27.31it/s]
│ Pairs: 78%|███████▊ | 310/400 [00:08<00:02, 32.25it/s]
│ Pairs: 78%|███████▊ | 314/400 [00:08<00:02, 30.15it/s]
│ Pairs: 80%|███████▉ | 318/400 [00:09<00:02, 28.20it/s]
│ Pairs: 80%|████████ | 321/400 [00:09<00:02, 27.64it/s]
│ Pairs: 81%|████████ | 324/400 [00:09<00:02, 27.98it/s]
│ Pairs: 82%|████████▎ | 330/400 [00:09<00:02, 32.86it/s]
│ Pairs: 84%|████████▎ | 334/400 [00:09<00:02, 30.60it/s]
│ Pairs: 84%|████████▍ | 338/400 [00:09<00:02, 28.66it/s]
│ Pairs: 85%|████████▌ | 341/400 [00:09<00:02, 27.93it/s]
│ Pairs: 86%|████████▋ | 345/400 [00:09<00:01, 29.03it/s]
│ Pairs: 88%|████████▊ | 350/400 [00:10<00:01, 33.47it/s]
│ Pairs: 88%|████████▊ | 354/400 [00:10<00:01, 31.27it/s]
│ Pairs: 90%|████████▉ | 358/400 [00:10<00:01, 29.85it/s]
│ Pairs: 90%|█████████ | 362/400 [00:10<00:01, 28.91it/s]
│ Pairs: 92%|█████████▏| 366/400 [00:10<00:01, 30.81it/s]
│ Pairs: 93%|█████████▎| 371/400 [00:10<00:00, 33.06it/s]
│ Pairs: 94%|█████████▍| 375/400 [00:10<00:00, 31.17it/s]
│ Pairs: 95%|█████████▍| 379/400 [00:11<00:00, 30.01it/s]
│ Pairs: 96%|█████████▌| 383/400 [00:11<00:00, 28.41it/s]
│ Pairs: 97%|█████████▋| 388/400 [00:11<00:00, 33.28it/s]
│ Pairs: 98%|█████████▊| 392/400 [00:11<00:00, 30.29it/s]
│ Pairs: 99%|█████████▉| 396/400 [00:11<00:00, 28.42it/s]
│ Pairs: 100%|█████████▉| 399/400 [00:11<00:00, 27.73it/s]
│ Pairs: 100%|██████████| 400/400 [00:11<00:00, 33.82it/s]
│
└─ Done (20 sequences · 11.83s)
Total running time of the script: (0 minutes 11.941 seconds)