def fit(*, identity: Literal["hash", "category"], seed: int, steps: int | None, accelerator: str) -> rf.Model:
"""Train one representation on the same identity pairs."""
lit.seed_everything(seed, workers=True)
model = rf.Model(
name="identity",
d_model=48,
n_layers=2,
n_heads=4,
batch_size=128,
left_id=rf.Hash(n_hashes=4) if identity == "hash" else rf.Category(size=8192, p_unavailable=0.0),
right_id=rf.Hash(n_hashes=4) if identity == "hash" else rf.Category(size=8192, p_unavailable=0.0),
equal=rf.Boolean(mask=True),
)
model.optimizer = lambda module: torch.optim.AdamW(module.parameters(), lr=3e-3)
data = rf.SyntheticDataModule(
model=model,
train=partial(records, rows=4096, seed=seed + 1, namespace="train"),
validate=partial(records, rows=1024, seed=seed + 2, namespace="validate"),
seed=seed,
)
trainer = lit.Trainer(
accelerator=accelerator,
max_epochs=20,
max_steps=steps if steps is not None else -1,
logger=False,
enable_progress_bar=False,
enable_model_summary=False,
enable_checkpointing=False,
deterministic=True,
)
trainer.fit(model=model, datamodule=data)
return model
def score(model: rf.Model, records: Callable[[], Iterator[dict]], accelerator: str) -> float:
"""Evaluate Boolean AUC without updating the vocabulary."""
data = rf.SyntheticDataModule(model=model, test=records)
trainer = lit.Trainer(
accelerator=accelerator,
logger=False,
enable_progress_bar=False,
enable_model_summary=False,
enable_checkpointing=False,
deterministic=True,
)
metrics = trainer.test(model=model, datamodule=data, verbose=False)[0]
return float(metrics["identity.equal/test.auc.content"])
def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
test = partial(records, rows=4096, seed=seed + 3, namespace="test")
control = partial(records, rows=4096, seed=seed + 3, namespace="test", shuffle_targets=True)
hash_model = fit(identity="hash", seed=seed, steps=steps, accelerator=accelerator)
category_model = fit(identity="category", seed=seed, steps=steps, accelerator=accelerator)
equality_auc = score(hash_model, test, accelerator)
control_auc = score(hash_model, control, accelerator)
category_oov_auc = score(category_model, test, accelerator)
gap = equality_auc - control_auc
return {
"hash_auc": equality_auc,
"shuffled_auc": control_auc,
"category_oov_auc": category_oov_auc,
"auc_gap": gap,
}, {
"Unseen Hash equality AUC is at least 0.95": equality_auc >= 0.95,
"Shuffled labels remain between 0.42 and 0.58 AUC": 0.42 <= control_auc <= 0.58,
"Category OOV AUC is at most 0.65": category_oov_auc <= 0.65,
"Hash exceeds shuffled labels by at least 0.35 AUC": gap >= 0.35,
}