Example of resampling when one class is underrepresented in the dataset.
import ROOT
import torch
seed = 42
torch.manual_seed(seed)
def make_df(b1_expr, num_events):
df_major = make_df("(int) 2 * rdfentry_", 100000)
df_minor = make_df("(int) 2 * rdfentry_ + 1", 1000)
batch_size = 256
num_epochs = 20
loss_fn = torch.nn.BCEWithLogitsLoss()
def train_model(model, optimizer, dataloader):
train, val = dataloader.train_test_split(test_size=0.2)
for _ in range(num_epochs):
train_correct = 0
train_total = 0
train_losses = []
model.train()
for X, y in train.as_torch():
optimizer.zero_grad()
outputs = model(X)
loss = loss_fn(outputs, y)
loss.backward()
optimizer.step()
preds = (outputs > 0.5).float()
train_correct += (preds == y).
sum().item()
train_total += y.size(0)
train_losses.append(loss.item())
print(
f"Training => Accuracy: {int(train_correct / train_total * 100000) / 100000}; Loss: {int(sum(train_losses) / len(train_losses) * 100000) / 100000}"
)
val_losses = []
val_correct = 0
val_total = 0
for X, y in val.as_torch():
with torch.no_grad():
outputs = model(X)
loss = loss_fn(outputs, y)
preds = (outputs > 0.5).float()
val_correct += (preds == y).
sum().item()
val_total += y.size(0)
val_losses.append(loss.item())
print(
f"Validation => Accuracy: {int(val_correct / val_total * 100000) / 100000}; Loss: {int(sum(val_losses) / len(val_losses) * 100000) / 100000}\n"
)
dl_oversampled = ROOT.Experimental.ML.RDataLoader(
[df_major, df_minor],
batch_size=batch_size,
target="b2",
set_seed=seed,
load_eager=True,
sampling_type="oversampling",
sampling_ratio=0.1,
)
oversampling_model = torch.nn.Linear(1, 1)
oversampling_optimizer = torch.optim.Adam(oversampling_model.parameters())
print("Training with oversampling:")
train_model(oversampling_model, oversampling_optimizer, dl_oversampled)
RInterface< Proxied > Define(std::string_view name, F expression, const ColumnNames_t &columns={})
Define a new column.
ROOT's RDataFrame offers a modern, high-level interface for analysis of data stored in TTree ,...
static uint64_t sum(uint64_t i)
Training with oversampling:
Training => Accuracy: 0.09089; Loss: 54054.4592
Training => Accuracy: 0.09089; Loss: 22879.47853
Training => Accuracy: 0.71861; Loss: 852.54963
Training => Accuracy: 0.9091; Loss: 1.21655
Training => Accuracy: 0.9091; Loss: 1.0924
Training => Accuracy: 0.9091; Loss: 0.938
Training => Accuracy: 0.9091; Loss: 0.74917
Training => Accuracy: 0.9091; Loss: 0.52145
Training => Accuracy: 0.90906; Loss: 0.2629
Training => Accuracy: 0.90798; Loss: 0.09193
Training => Accuracy: 0.90683; Loss: 0.07532
Training => Accuracy: 0.90621; Loss: 0.07395
Training => Accuracy: 0.90539; Loss: 0.07244
Training => Accuracy: 0.90465; Loss: 0.07081
Training => Accuracy: 0.90397; Loss: 0.06906
Training => Accuracy: 0.90324; Loss: 0.06723
Training => Accuracy: 0.9283; Loss: 0.06534
Training => Accuracy: 0.99265; Loss: 0.06342
Training => Accuracy: 0.99201; Loss: 0.06152
Training => Accuracy: 0.99136; Loss: 0.05965
Validation => Accuracy: 0.99016; Loss: 0.05549
- Author
- Jonah Ascoli
Definition in file ml_dataloader_resampling.py.