Feature: OCR for item drops and item description boxes (#570)

* initial

* bugfixes. now stable & testing

* narrow ROI to search for 'CLICK TO' tooltip
This commit is contained in:
mgleed
2022-02-21 19:54:29 -05:00
committed by GitHub
parent 8fa729ddab
commit c8eb88f672
16 changed files with 237 additions and 25 deletions
Binary file not shown.

After

Width:  |  Height:  |  Size: 2.1 KiB

+1
View File
@@ -1,6 +1,7 @@
SHIFLD, SHIELD
SPFAR, SPEAR
GLOVFS, GLOVES
GOLP, GOLD
TELEFORT, TELEPORT
TROPHV, TROPHY
CLAVMORE, CLAYMORE
1 SHIFLD SHIELD
2 SPFAR SPEAR
3 GLOVFS GLOVES
4 GOLP GOLD
5 TELEFORT TELEPORT
6 TROPHV TROPHY
7 CLAVMORE CLAYMORE
+1
View File
@@ -2,6 +2,7 @@
; min and max hsv range (opencv format: h: [0-180], s: [0-255], v: [0, 255])
; h_min, s_min, v_min, h_max, s_max, v_max
black=0,0,0,180,255,15
black_descr=0,0,0,180,255,25
item_highlight=90,235,130,115,255,160
white=0,0,150,180,20,255
gray=0,0,90,180,20,130
+1
View File
@@ -255,3 +255,4 @@ hwnd_window_title=
hwnd_window_process=D2R\.exe
;If you want to control Hyper-V window from host use 0,51 here
window_client_area_offset=0,0
use_ocr=0
+2 -1
View File
@@ -316,7 +316,8 @@ class Config:
"graphic_debugger_key": self._select_val("advanced_options", "graphic_debugger_key"),
"hwnd_window_title": _default_iff(Config()._select_val("advanced_options", "hwnd_window_title"), ''),
"hwnd_window_process": _default_iff(Config()._select_val("advanced_options", "hwnd_window_process"), ''),
"window_client_area_offset": tuple(map(int, Config()._select_val("advanced_options", "window_client_area_offset").split(",")))
"window_client_area_offset": tuple(map(int, Config()._select_val("advanced_options", "window_client_area_offset").split(","))),
"use_ocr": bool(int(self._select_val("advanced_options", "use_ocr"))),
}
self.items = {}
+2 -2
View File
@@ -41,7 +41,7 @@ class GameStats:
if self._location not in self._location_stats:
self._location_stats[self._location] = { "items": [], "deaths": 0, "chickens": 0, "merc_deaths": 0, "failed_runs": 0 }
def log_item_keep(self, item_name: str, send_message: bool, img: np.ndarray):
def log_item_keep(self, item_name: str, send_message: bool, img: np.ndarray, ocr_text: str = None):
Logger.debug(f"Stashed and logged: {item_name}")
filtered_items = ["_potion", "misc_gold"]
if self._location is not None and not any(substring in item_name for substring in filtered_items):
@@ -49,7 +49,7 @@ class GameStats:
self._location_stats["totals"]["items"] += 1
if send_message:
self._messenger.send_item(item_name, img, self._location)
self._messenger.send_item(item_name, img, self._location, ocr_text)
def log_death(self, img: str):
self._death_counter += 1
+70 -7
View File
@@ -4,20 +4,31 @@ from dataclasses import dataclass
import time
from utils.misc import color_filter, erode_to_black
from template_finder import TemplateFinder
from ocr import Ocr, OcrResult
from config import Config
from logger import Logger
# TODO: With OCR we can then add a "text" field to this class
@dataclass
class ItemText:
color_key: str = None
color: str = None
roi: list[int] = None
data: np.ndarray = None
ocr_result: OcrResult = None
clean_img: np.ndarray = None
def __getitem__(self, key):
return super().__getattribute__(key)
class ItemCropper:
def __init__(self):
self._ocr = Ocr()
self._gaus_filter = (19, 1)
self._expected_height_range = [round(num) for num in [x / 1.5 for x in [14, 40]]]
self._expected_width_range = [round(num) for num in [x / 1.5 for x in [60, 1280]]]
self._box_expected_width_range=[200, 900]
self._box_expected_height_range=[24, 710]
self._hud_mask = cv2.imread(f"assets/hud_mask.png", cv2.IMREAD_GRAYSCALE)
self._hud_mask = cv2.threshold(self._hud_mask, 1, 255, cv2.THRESH_BINARY)[1]
@@ -70,28 +81,80 @@ class ItemCropper:
max_idx = color_averages.index(max(color_averages))
if key == self._item_colors[max_idx]:
item_clusters.append(ItemText(
color_key = key,
color = key,
roi = [x, y, w, h],
data = cropped_item
data = cropped_item,
clean_img = cleaned_img[y:y+h, x:x+w]
))
debug_str += f" | cluster: {time.time() - start}"
# print(debug_str)
if Config().advanced_options["use_ocr"]:
cluster_images = [ key["clean_img"] for key in item_clusters ]
results = self._ocr.image_to_text(cluster_images, model = "engd2r_inv_th_fast", psm = 7)
for count, cluster in enumerate(item_clusters):
setattr(cluster, "ocr_result", results[count])
return item_clusters
def crop_item_descr(self, inp_img: np.ndarray, all_results: bool = False, inventory_side: str = "right") -> ItemText:
"""
Crops visible item description boxes / tooltips
:inp_img: image from hover over item of interest.
:param all_results: whether to return all possible results (True) or the first result (False)
:inventory_side: enter either "left" for stash/vendor region or "right" for user inventory region
"""
results=[]
black_mask, _ = color_filter(inp_img, Config().colors["black"])
contours = cv2.findContours(black_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
contours = contours[0] if len(contours) == 2 else contours[1]
for cntr in contours:
x, y, w, h = cv2.boundingRect(cntr)
cropped_item = inp_img[y:y+h, x:x+w]
avg = np.average(cv2.cvtColor(cropped_item, cv2.COLOR_BGR2GRAY))
mostly_dark = True if 0 < avg < 20 else False
contains_black = True if np.min(cropped_item) < 14 else False
contains_white = True if np.max(cropped_item) > 250 else False
contains_orange = False
if not contains_white:
#check for orange (like key of destruction, etc.)
orange_mask, _ = color_filter(cropped_item, Config().colors["orange"])
contains_orange = np.min(orange_mask) > 0
expected_height = True if (self._box_expected_height_range[0] < h < self._box_expected_height_range[1]) else False
expected_width = True if (self._box_expected_width_range[0] < w < self._box_expected_width_range[1]) else False
box2 = Config().ui_roi[f"{inventory_side}_inventory"]
overlaps_inventory = False if (x+w<box2[0] or box2[0]+box2[2]<x or y+h+28+10<box2[1] or box2[1]+box2[3]<y) else True # padded height because footer isn't included in contour
if contains_black and (contains_white or contains_orange) and mostly_dark and expected_height and expected_width and overlaps_inventory:
footer_height_max = (720 - (y + h)) if (y + h + 35) > 720 else 35
found_footer = TemplateFinder().search(["TO_TOOLTIP"], inp_img, threshold=0.8, roi=[x, y+h, w, footer_height_max]).valid
if found_footer:
ocr_result = None
if Config().advanced_options["use_ocr"]:
ocr_result = self._ocr.image_to_text(cropped_item, psm=6)[0]
results.append(ItemText(
color = "black",
roi = [x, y, w, h],
data = cropped_item,
ocr_result = ocr_result
))
if not all_results:
break
return results
if __name__ == "__main__":
import keyboard
import os
from screen import grab
from template_finder import TemplateFinder
keyboard.add_hotkey('f12', lambda: os._exit(1))
cropper = ItemCropper()
while 1:
img = grab().copy()
res = cropper.crop(img)
for cluster in res:
x, y, w, h = cluster.roi
cv2.rectangle(img, (x, y), (x+w, y+h), (0, 255, 0), 1)
results = cropper.crop_item_descr(img, all_results=True, ocr=False)
for res in results:
if res["color"]:
x, y, w, h = res.roi
cv2.rectangle(img, (x, y), (x+w, y+h), (0, 255, 0), 1)
Logger.debug(f"{res.ocr_result['text']}")
cv2.imshow("res", img)
cv2.waitKey(1)
+8 -1
View File
@@ -9,6 +9,8 @@ import math
from config import Config
from utils.misc import color_filter, cut_roi
from item import ItemCropper
from template_finder import TemplateFinder
from ocr import OcrResult
@dataclass
@@ -24,6 +26,8 @@ class Item:
score: float = -1.0
dist: float = -1.0
roi: list[int] = None
color: str = None
ocr_result: OcrResult = None
def __getitem__(self, key):
return super().__getattribute__(key)
@@ -119,6 +123,8 @@ class ItemFinder:
item.roi = [max_loc[0] + x, max_loc[1] + y, template.data.shape[1], template.data.shape[0]]
center_abs = (item.center[0] - (inp_img.shape[1] // 2), item.center[1] - (inp_img.shape[0] // 2))
item.dist = math.dist(center_abs, (0, 0))
item.ocr_result = cluster.ocr_result
item.color = cluster.color
if item is not None and self._items_to_pick[item.name].pickit_type:
item_list.append(item)
elapsed = time.time() - start
@@ -130,6 +136,7 @@ class ItemFinder:
if __name__ == "__main__":
from screen import grab
from config import Config
item_finder = ItemFinder()
while 1:
# img = cv2.imread("")
@@ -139,7 +146,7 @@ if __name__ == "__main__":
# print(item.name + " " + str(item.score))
cv2.circle(img, item.center, 5, (255, 0, 255), thickness=3)
cv2.rectangle(img, item.roi[:2], (item.roi[0] + item.roi[2], item.roi[1] + item.roi[3]), (0, 0, 255), 1)
# cv2.putText(img, item.name, item.center, cv2.FONT_HERSHEY_SIMPLEX, 0.8, (255, 255, 255), 1, cv2.LINE_AA)
cv2.putText(img, item.ocr_result["text"], item.center, cv2.FONT_HERSHEY_SIMPLEX, 0.5, (255, 255, 255), 1, cv2.LINE_AA)
# img = cv2.resize(img, None, fx=0.5, fy=0.5)
cv2.imshow('test', img)
cv2.waitKey(1)
+23
View File
@@ -40,6 +40,8 @@ class PickIt:
curr_item_to_pick: Item = None
same_item_timer = None
did_force_move = False
done_ocr=False
while not time_out:
if (time.time() - start) > 28:
time_out = True
@@ -48,6 +50,23 @@ class PickIt:
img = grab()
item_list = self._item_finder.search(img)
if Config().advanced_options["use_ocr"] and not done_ocr:
timestamp = time.strftime("%Y%m%d_%H%M%S")
for cnt, item in enumerate(item_list):
for cnt2, x in enumerate(item.ocr_result['word_confidences']):
found_low_confidence = False
if x <= 88:
try:
Logger.debug(f"Low confidence word #{cnt2}: {item.ocr_result['original_text'].split()[cnt2]} -> {item.ocr_result['text'].split()[cnt2]}, Conf: {x}, save screenshot")
found_low_confidence = True
except: pass
if found_low_confidence and Config().general["loot_screenshots"]:
cv2.imwrite(f"./loot_screenshots/ocr_drop_{timestamp}_{cnt}_o.png", item.ocr_result['original_img'])
cv2.imwrite(f"./loot_screenshots/ocr_drop_{timestamp}_{cnt}_n.png", item.ocr_result['processed_img'])
with open(f"./loot_screenshots/ocr_drop_{timestamp}_{cnt}_o.gt.txt", 'w') as f:
f.write(item.ocr_result['text'])
done_ocr = True
# Check if we need to pick up certain pots more pots
need_pots = belt.get_pot_needs()
if need_pots["mana"] <= 0:
@@ -110,6 +129,10 @@ class PickIt:
# no need to stash potions, scrolls, or gold
if "potion" not in closest_item.name and "tp_scroll" != closest_item.name and "misc_gold" not in closest_item.name:
found_items = True
if Config().advanced_options["use_ocr"]:
for item in item_list:
Logger.debug(f"OCR DROP: Name: {item.ocr_result['text']}, Conf: {item.ocr_result['word_confidences']}")
prev_cast_start = char.pick_up_item((x_m, y_m), item_name=closest_item.name, prev_cast_start=prev_cast_start)
if not char.capabilities.can_teleport_natively:
+3 -2
View File
@@ -18,7 +18,7 @@ class DiscordEmbeds(GenericApi):
except InvalidArgument:
Logger.warning(f"Your custom_message_hook URL {Config().general['custom_message_hook']} is invalid, Discord updates will not be sent")
def send_item(self, item: str, image: np.ndarray, location: str):
def send_item(self, item: str, image: np.ndarray, location: str, ocr_text: str = None):
imgName = item.replace('_', '-')
_, w, _ = image.shape
@@ -32,6 +32,7 @@ class DiscordEmbeds(GenericApi):
)
e.set_thumbnail(url=f"{self._psnURL}41L6bd712.png")
e.set_image(url=f"attachment://{imgName}.png")
e.add_field(name="OCR Text", value=f"{ocr_text}", inline=False)
self._send_embed(e, file)
def send_death(self, location, image_path):
@@ -96,7 +97,7 @@ class DiscordEmbeds(GenericApi):
return Color.blue()
def _add_file(self, image_path, image_name):
try:
try:
return discord.File(image_path, filename=image_name)
except:
traceback.print_exc()
+5 -5
View File
@@ -6,18 +6,18 @@ import requests
class GenericApi:
def send_item(self, item: str, image: np.ndarray, location: str):
def send_item(self, item: str, image: np.ndarray, location: str, ocr_text: str = None):
msg = f"Found {item} at {location}"
self._send(msg)
def send_death(self, location: str, image_path: str = None):
msg = f"You have died at {location}"
self._send(msg)
def send_chicken(self, location: str, image_path: str = None):
msg = f"You have chickened at {location}"
self._send(msg)
def send_gold(self):
msg = f"All stash tabs and character are full of gold, turn of gold pickup"
self._send(msg)
@@ -31,7 +31,7 @@ class GenericApi:
def _send(self, msg: str):
msg = f"{Config().general['name']}: {msg}"
url = Config().general['custom_message_hook']
if not url:
return
+2 -2
View File
@@ -15,8 +15,8 @@ class Messenger:
else:
self._message_api = None
def send_item(self, item: str, image: np.ndarray, location: str):
self._message_api.send_item(item, image, location)
def send_item(self, item: str, image: np.ndarray, location: str, ocr_text: str = None):
self._message_api.send_item(item, image, location, ocr_text)
def send_death(self, location: str, image_path: str = None):
self._message_api.send_death(location, image_path)
+22 -4
View File
@@ -1,3 +1,4 @@
from concurrent.futures import process
from fileinput import close
from tesserocr import PyTessBaseAPI, PSM, OEM
import numpy as np
@@ -119,10 +120,9 @@ class Ocr:
if crop_pad:
processed_img = self._crop_pad(processed_img)
image_is_binary = (image.shape[2] if len(image.shape) == 3 else 1) == 1 and image.dtype == bool
if image_is_binary:
if threshold:
processed_img = cv2.cvtColor(processed_img, cv2.COLOR_BGR2GRAY)
processed_img = cv2.threshold(processed_img, threshold, 255, cv2.THRESH_BINARY)[1]
if not image_is_binary and threshold:
processed_img = cv2.cvtColor(processed_img, cv2.COLOR_BGR2GRAY)
processed_img = cv2.threshold(processed_img, threshold, 255, cv2.THRESH_BINARY)[1]
if invert:
if threshold or image_is_binary:
processed_img = cv2.bitwise_not(processed_img)
@@ -239,6 +239,24 @@ class Ocr:
continue
break
# case: a solitary S; e.g., " 1 TO S DEFENSE"
cnt=0
while True:
cnt += 1
if cnt >30:
Logger.error(f"Error ' S ' -> ' 5 ' on {ocr_output}")
break
if " S " in text:
text = text.replace(" S ", " 5 ")
continue
elif ' I\n' in text:
text = text.replace(' S\n', ' 5\n')
continue
elif '\nI ' in text:
text = text.replace('\nS ', '\n5 ')
continue
break
# case: consecutive I's; e.g., "DEFENSE: II"
repeat=False
cnt=0
-1
View File
@@ -1,7 +1,6 @@
import cv2
import threading
from copy import deepcopy
from item.item_finder import Template
from screen import convert_screen_to_monitor, grab
from typing import Union
from dataclasses import dataclass
+23
View File
@@ -16,6 +16,7 @@ from ui_components import stash
from ui_components.stash import gold_full
from ui.ui_manager import detect_screen_object, messenger, game_stats, wait_for_screen_object, ScreenObjects
from messages import Messenger
from item import ItemCropper
messanger = Messenger()
@@ -180,6 +181,28 @@ def keep_item(item_finder: ItemFinder, img: np.ndarray, do_logging: bool = True)
"""
wait(0.2, 0.3)
_, w, _ = img.shape
if Config().advanced_options["use_ocr"]:
item_box = ItemCropper().crop_item_descr(inp_img=img)
if item_box:
item_box = item_box[0]
Logger.debug(f"OCR ITEM DESCR: Mean conf: {item_box.ocr_result.mean_confidence}")
for i, line in enumerate(list(filter(None, item_box.ocr_result.text.splitlines()))):
Logger.debug(f"OCR LINE{i}: {line}")
if Config().general["loot_screenshots"]:
timestamp = time.strftime("%Y%m%d_%H%M%S")
found_low_confidence = False
for cnt, x in enumerate(item_box.ocr_result['word_confidences']):
if x <= 88:
try:
Logger.debug(f"Low confidence word #{cnt}: {item_box.ocr_result['original_text'].split()[cnt]} -> {item_box.ocr_result['text'].split()[cnt]}, Conf: {x}, save screenshot")
found_low_confidence = True
except: pass
if found_low_confidence:
cv2.imwrite(f"./loot_screenshots/ocr_box_{timestamp}_o.png", item_box.ocr_result['original_img'])
cv2.imwrite(f"./loot_screenshots/ocr_box_{timestamp}_n.png", item_box.ocr_result['processed_img'])
img = img[:, (w//2):,:]
original_list = item_finder.search(img)
filtered_list = []
+74
View File
@@ -0,0 +1,74 @@
from screen import Screen
import cv2
from config import Config
from utils.misc import cut_roi
import mouse
import keyboard
import os
import time
from screen import grab, convert_monitor_to_screen
class GenOcrTruth:
def __init__(self):
if not os.path.exists("generated"):
os.system("mkdir generated")
os.system(f"cd generated && mkdir ground-truth")
self._half_width = Config().ui_pos["screen_width"] // 2
self._half_height = Config().ui_pos["screen_height"] // 2
self._upper_left = None
def hook(self, e):
if e.event_type == "down":
if e.name == "f12":
os._exit(1)
img = grab()
loc_monitor = mouse.get_position()
loc_screen = convert_monitor_to_screen(loc_monitor)
if e.name == "f8":
# start template
if self._upper_left is None:
self._upper_left = loc_screen
print(f"stored upper left: {self._upper_left}")
print("Select bottom-right corner of template to create and press f8")
return
# finish template
else:
bottom_right = loc_screen
if bottom_right == self._upper_left:
print(f"roi failed, try again")
self._upper_left = None
return
print(f"stored bottom_right: {bottom_right}")
try:
width = (bottom_right[0] - self._upper_left[0])
height = (bottom_right[1] - self._upper_left[1])
template_img = cut_roi(img, [*self._upper_left, width, height])
basename = f"generated/ground-truth/{time.strftime('%Y%m%d_%H%M%S')}"
cv2.imshow(time.strftime('%Y%m%d_%H%M%S'), template_img)
cv2.waitKey(1)
print(f"new template: {basename} = ")
print(f"Enter true text:")
truth = input()
if truth:
with open(f"{basename}.gt.txt", 'w') as f:
f.write(truth)
cv2.imwrite(f"{basename}.png", template_img)
cv2.destroyAllWindows()
cv2.waitKey(1)
print(f"saved {basename}")
else:
print(f"skipped {basename}")
except:
print(f"template save failed, try again")
self._upper_left = None
print("Select top-left corner of template to create and press f8 | F12 to exit")
return
else:
return
if __name__ == "__main__":
keyboard.add_hotkey('f12', lambda: print('Force Exit (f12)') or os._exit(1))
recorder = GenOcrTruth()
print("Select top-left corner of template to create and press f8 | F12 to exit")
keyboard.hook(recorder.hook, suppress=True)
while True: pass