Views
No views yet
1base_model = ViTModel.from_pretrained("google/vit-base-patch16-224-in21k")
2
3class ViTForRegression(nn.Module):
4 def __init__(self, base_model, num_outputs=2):
5 super(ViTForRegression, self).__init__()
6 self.base_model = base_model
7 hidden_size = base_model.config.hidden_size
8 self.regression_head = nn.Linear(hidden_size, num_outputs)
9
10 def forward(self, pixel_values):
11 outputs = self.base_model(pixel_values=pixel_values)
12 pooler_output = outputs.pooler_output
13 predictions = self.regression_head(pooler_output)
14 return predictions
15
16model = ViTForRegression(base_model).to(device)dataset_test = load_dataset("gydou/released_img")1lat_std = 0.0006914493505038013
2lon_std = 0.0006539239061573955
3lat_mean = 39.9517411499467
4lon_mean = -75.19143213125122