Files
vector_search/search_me.py
T
2026-07-09 10:26:27 -04:00

110 lines
3.4 KiB
Python

do_load = True
import requests
from pymilvus import MilvusClient, DataType
from pymilvus.client.types import LoadState
from datetime import datetime, timedelta
import numpy as np
import traceback
from bottle import route, run, template, request, debug
import time
import prettyprinter
TIMEOUT=10
collection_name = "nuggets_{camera}_so400m_siglip2"
client = MilvusClient(
uri="http://localhost:19530"
)
from bottle import route, run, template
@route('/get_collection_count')
def get_collection_count():
coll_count = dict()
for c in sorted(client.list_collections()):
coll_count[c] = client.get_collection_stats(c)['row_count']
return coll_count
@route('/get_text_match')
def get_matches():
query = request.query.get('query','A large bird eating corn')
cameras = request.query.get('cameras','sidefeeder')
num_videos = int(request.query.get('num_videos',5))
cams = set(cameras.split(","))
max_age = int(request.query.get('age',5));
max_date = datetime.now()
min_date = max_date - timedelta(days=(max_age-1))
days_step = (max_date- min_date).days
day_strs = list()
for x in range(days_step):
# pass
day_strs.append( (min_date + timedelta(days=x)).strftime('%Y%m%d') )
# if days_step > 0:
# day_strs.append( (min_date + timedelta(days=0)).strftime('%Y%m%d') )
day_strs.append( max_date.strftime('%Y%m%d'))
str_insert = ','.join([f'"{x}"'for x in day_strs])
filter_string = 'date in [' + str_insert + ']'
if max_age == 0:
filter_string = ''
vec_form = requests.get('http://192.168.1.242:53004/encode',params={'query':query}).json()['vector'][0]
vec_search = np.asarray(vec_form).astype(np.float16)
print('Filter string:',filter_string)
all_results=list()
try:
error = ''
if True:
for cam in cams:
col_name = collection_name.format(camera=cam)
for i in [num_videos]:
results = client.search(collection_name = col_name,
consistency_level="Eventually",
filter=filter_string,
data = [vec_search],
limit=i,
search_params={'metric_type':'COSINE', 'params':{}},
output_fields=['filepath','frame_number']
)
all_results.extend(results[0])
except Exception as e:
error = traceback.format_exc()
print(error)
results = []
def linux_to_win_path(form):
form = form.replace('/srv','file://192.168.1.242/thebears/Videos/merged/')
return form
def normalize_to_merged(path):
if not path.startswith('/'):
path= '/mergedfs/ftp/'+path
return path
resul = list()
all_results = sorted(all_results, key=lambda x: x['distance'])
for x in all_results:
pload =dict( x['entity'])
pload['filepath'] = normalize_to_merged(pload['filepath'])
pload['score'] = x['distance']
pload['winpath'] = linux_to_win_path(pload['filepath'])
pload['frame'] = pload['frame_number']
resul.append(pload)
return_this = {'query':query,'num_videos':num_videos,'results':resul,'error':error}
return return_this
debug(True)
run(host='0.0.0.0', port=53003, server='bjoern')