model = CustomViTRegressor(should_load_from_disk=False)
model.update_model_from_checkpoint(2)
custom_data_collator = custom_data_collator_function(ViTImageProcessor())
train_loader = DataLoader(dataset["train"], batch_size=32, collate_fn=custom_data_collator)
model = model.to(device)
def custom_data_collator_function(processor):
def return_func(batch):
images = [item["image"] for item in batch]
inputs = processor(images, return_tensors="pt")
targets = torch.tensor([[
item["Start"],
item["A"],
item["B"],
item["X"],
item["Y"],
item["Z"],
item["DPadUp"],
item["DPadDown"],
item["DPadLeft"],
item["DPadRight"],
item["L"],
item["R"],
item["LPressure"] / 255,
item["RPressure"] / 255,
item["XAxis"] / 255,
item["YAxis"] / 255,
item["CXAxis"] / 255,
item["CYAxis"] / 255] for item in batch], dtype=torch.float32)
return inputs, targets
return return_func
def is_interesting_target(target):
if target[0] == 1 or \
target[1] == 1 or \
target[2] == 1 or \
target[3] == 1 or \
target[4] == 1 or \
target[5] == 1 or \
target[6] == 1 or \
target[7] == 1 or \
target[8] == 1 or \
target[9] == 1 or \
target[10] == 1 or \
target[11] == 1:
return 1
if target[12] >= THRESHOLD or \
target[13] >= THRESHOLD or \
abs(target[14] - 127.5) >= THRESHOLD or \
abs(target[15] - 127.5) >= THRESHOLD or \
abs(target[16] - 127.5) >= THRESHOLD or \
abs(target[17] - 127.5) >= THRESHOLD:
return 1
return 0