summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorHunternif <hunternif@gmail.com>2021-08-31 06:10:01 +0100
committerHunternif <hunternif@gmail.com>2021-08-31 06:14:47 +0100
commit2a149935d14936a158e5ac022fba10b65c9c02dc (patch)
treebf7c46065a007b4d0e6a891e0ada271cbff26492
parent0ec6f30125cda1227104280f1a61fe6ca1cff59a (diff)
server: simplify similar post logic for better performance
-rw-r--r--server/szurubooru/func/similar.py31
1 files changed, 15 insertions, 16 deletions
diff --git a/server/szurubooru/func/similar.py b/server/szurubooru/func/similar.py
index 8247de8..77d7805 100644
--- a/server/szurubooru/func/similar.py
+++ b/server/szurubooru/func/similar.py
@@ -1,3 +1,4 @@
+from math import ceil
from queue import Queue
from typing import List
@@ -6,26 +7,19 @@ from szurubooru import model, search
_search_executor_config = search.configs.PostSearchConfig()
_search_executor = search.Executor(_search_executor_config)
-_max_removals = 3
-
def find_similar_posts(source_post: model.Post, limit: int) -> List[model.Post]:
results = []
- queue = Queue() # contains lists of tags to search
- queue.put(source_post.tags)
- source_tag_count = len(source_post.tags)
-
- while not queue.empty():
- # put follow-up searches on the queue
- last_tags = queue.get()
- tag_count = len(last_tags)
- if tag_count > 1 and tag_count > source_tag_count - _max_removals:
- for removed_tag in last_tags:
- next_search = list(filter(lambda t: t != removed_tag, last_tags))
- queue.put(next_search)
+ # Sort tags in order of increasing post count, i.e. least to most popular
+ # This will help yield results quicker
+ source_tags = sorted(source_post.tags, key=lambda t: t.post_count)
+ source_tag_count = len(source_tags)
+ max_removals = ceil(source_tag_count / 2) # remove at most 50% of tags
+ tags = source_tags
+ for x in range(max_removals + 1):
# prepare the current search, remove known results
- query = ' '.join([t.first_name for t in last_tags])
+ query = ' '.join([t.first_name for t in tags])
query += ' -id:%d' % source_post.post_id
for r in results:
query += ' -id:%d' % r.post_id
@@ -37,6 +31,11 @@ def find_similar_posts(source_post: model.Post, limit: int) -> List[model.Post]:
for p in posts:
results.append(p)
if len(results) >= limit:
- return results
+ break
+
+ # remove the least popular tag
+ if len(tags) <= 1:
+ break
+ tags = tags[1:]
return results