Description
This model demonstrates the capabilities of a CNN for classifying screenshots as game/reality, using Minecraft as an example:
0 class - Minecraft.
1 class - Not Minecraft.
Main Idea
The hypothesis was that the game Minecraft has distinct 3D voxel graphics, which would allow the model to distinguish photos from the game from photos taken in real life.
The model was trained on 4,500 screenshot photos from the game and approximately 3,800 photos from real life and other games.
Architecture
The model is a 4‑layer and 3-layer convolutional networks and two linear MLP layers for classification.
Tests
The first test was conducted on vanilla Minecraft/real photos.
| Layer count | F1 | Accuracy |
|---|---|---|
| 3 | 0.94 | 0.94 |
| 4 | 0.96 | 0.96 |
The second test was conducted on Minecraft (shaders, texture packs, and vanilla version)/ real photo, photos of voxel games (3D cube world, Dragon Quest Builders).
| Layer count | F1 | Accuracy |
|---|---|---|
| 3 | 0.83 | 0.83 |
| 4 | 0.82 | 0.82 |
The test results show that the Minecraft/Reality task is simpler for CNN than Minecraft/Similar 3D voxel games. The hypothesis was confirmed for simple examples. For more complex tasks, other methods are required, including increasing the dataset size, changing the problem statement, and providing a better variety of training data.
Inferense example
from transformers import AutoModel
import torch
from torchvision import transforms
from PIL import Image
from transformers import AutoConfig, AutoModelForImageClassification
def open_image(image_path):
image = Image.open(image_path).convert('RGB')
image_ten = transform(image)
return image_ten
def get_prediction(model, image):
model.eval()
with torch.no_grad():
pred = model(image)
print(pred)
out = torch.softmax(pred, dim=1)
print(out)
out = torch.argmax(out, dim=1)
print(out)
return out
repo_id = 'Neweret/minecraft_screenshot_classifier'
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = AutoModelForImageClassification.from_pretrained(
repo_id,
trust_remote_code = True
).to(device)
transform = transforms.Compose([
transforms.Resize((224,224)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
print('Write the image path: ')
image_path = input().strip()
try:
image_ten = open_image(image_path).to(device).unsqueeze(0)
except AttributeError:
print('Error: Incorrect value type!')
exit()
except FileNotFoundError:
print('Error: File not exist!')
exit()
pred = get_prediction(model, image_ten)
if pred == 0:
print('It`s a Minecraft')
else:
print('It`s not Minecraft')
- Downloads last month
- -
