Google Research ha lanciato TabFM, un modello fondamentale per dati tabulari che può svolgere classificazione e regressione senza un addestramento specifico. Il modello si basa sull'apprendimento contestuale (in-context learning) e un'architettura ibrida, ottenendo predizioni direttamente da un singolo passaggio. Questo permette di effettuare previsioni su dati non visti mantenendo alta efficienza. TabFM è attualmente disponibile su Hugging Face e GitHub, e verrà integrato in Google BigQuery tramite il comando AI.PREDICT.

Che cos'è TabFM?

I dati tabulari rappresentano la spina dorsale delle infrastrutture aziendali. Molti compiti importanti, come la previsione di clienti in fuga o la rilevazione di frodi finanziarie, risiedono in tabelle strutturate. Negli ultimi anni, i metodi basati su alberi, come XGBoost, AdaBoost e random forests, hanno dominato questo ambito.

Questo tipo di approccio richiede comunque un'ottimizzazione manuale estesa. Data scientist passano ore per calibrare e raffinare i parametri. TabFM mira esattamente a questa criticità: offre un'alternativa in grado di effettuare previsioni su dati non visti direttamente, mantenendo elevata efficacia.

Come funziona

Concettualmente, TabFM opera in modo simile a un linguaggio modello (LM) quando genera previsioni. Usa l’apprendimento contestuale (in-context learning, ICL) per ottenere risultati da tabelle nuove, senza aggiornare i pesi interni. Il modello riceve l’intera tabella come singolo prompt, e sfrutta la relazione tra colonne e righe durante l'output.

Le tabelle a due dimensioni si distinguono per essere ordine e struttura flessibile. Cambiare l'ordine delle righe o delle colonne non ha alcuna rilevanza per il significato. Gli standard moderni dei modelli linguaggi non riescono a gestire questo contesto. Per bridgare questa differenza, TabFM integra modelli TabPFN e TabICL utilizzando un processo ibrido.

La sua architettura si basa su tre meccanismi:

    • Attenzione alternata righe-colonne: Il modello esegue attenzione multistrato su righe e colonne, permettendo di catturare relazioni complesse nel contesto.
    • Compressione delle righe: Ciascuna riga generata viene compressa in un singolo vettore denso.
    • Apprendimento contestuale: Un Transformer dedicato elabora i vettori, riducendo drasticamente i costi computazionali.

Allenamento su dati sintetici a grande scala

I modelli fondamentali richiedono dati ampi e diversificati. I dataset tabulari di alta qualità sono rari, soprattutto se aperti. Molti database industriali contengono informazioni sensibili, che non si possono utilizzare per un addestramento su larga scala.

TabFM è stato addestrato interamente su dati generati sinteticamente, usando metodi basati su modelli causali strutturati (SCMs). Questi dati simulano le relazioni complesse tipiche delle tabelle reali. Questo permette al modello di fare previsioni coerenti e significative anche su dati di contesto totalmente diversi.

Prestazioni e test

Il team di ricerca ha valutato TabFM su TabArena, una piattaforma con benchmark costante che valuta risultati con metriche Elo. Le prove coprono 38 dataset di classificazione e 13 dataset di regressione. La dimensione varia da 700 a 150,000 righe.

Sono state effettuate due configurazioni: TabFM e TabFM-Ensemble. La prima versione si usa direttamente, con risultati immediati. La seconda aggiunge funzioni incrociate e SVD, migliorando ulteriormente le performance.

TabFM supera regolarmente modelli supervisionati di alta qualità, come i modelli GBDT (Gradient Boosted Decision Trees) standard.

Iniziare con TabFM: installazione e codice

Per installare TabFM, si clona direttamente il repository. L'installazione base utilizza solo CPU con JAX. Per l’esecuzione GPU, vengono forniti plugin aggiuntivi di CUDA.

Requisiti principali:

    • Python 3.11 o versione successiva
    • Installazione di JAX e FLAX
    • Download automatico dei pesi pre-addestrati tramite Hugging Face

Codice esemplificativo

Predizione di rischio con TabFM

Di seguito è mostrato un esempio di utilizzo del modulo TabFMClassifier per effettuare classificazioni:


import numpy as np

import pandas as pd

from tabfm import tabfmv10_0

from tabfm import TabFMClassifier

Scarica i modelli pre-addestrati

model = tabfmv10_0.load()

Costruisci il classificatore

clf = TabFMClassifier(model=model)

Dataset di addestramento

X_train = pd.DataFrame({

"age": [25.0, 45.0, 35.0, 50.0],

"job": ["engineer", "manager", "engineer", "manager"],

"income": [80000, 120000, 90000, 130000]

})

ytrain = np.array(["lowrisk", "highrisk", "lowrisk", "high_risk"])

Dataset di test

X_test = pd.DataFrame({

"age": [30.0, 48.0],

"job": ["engineer", "manager"],

"income": [85000, 125000]

})

Esecuzione

clf.fit(Xtrain, ytrain)

predictions = clf.predict(X_test)

probabilities = clf.predictproba(Xtest)

print("Predictions:", predictions)

print("Class Probabilities:\n", probabilities)

L’implementazione non richiede un addestramento specifico su dati personali. La versione regressiva funziona nello stesso modo utilizzando TabFMRegressor.

Usi Pratici Con Esempi

Churn dei clienti

In casi come la fuga dei clienti, TabFM utilizza l'intero dataset per generare una previsione in una sola passata. Il contenuto del dataset include informazioni di clienti precedenti, contrassegnati come retenuti o smobilitati. Per ogni cliente nuovo, genera il livello di rischio senza ulteriore training.

Rischi di credito

Un'applicazione tipica include valutare nuovi clienti utilizzando informazioni di età, lavoro e reddito. TabFM genera valutazioni di rischio direttamente dagli stessi parametri, senza necessità di cicli di training supplementari.

Previsione del prezzo delle case

Per la regressione, ad esempio nella valutazione del costo immobiliare, le righe si basano su informazioni di metratura e quartiere. TabFM restituisce previsioni dirette per nuove proprietà