import sys rt_path = "/home/thebears/Source/task_runners/vision_v3/cuda_objdet_clip" sys.path.insert(0, rt_path) import torch cc = torch.cuda.get_device_properties(0) cc_str = str(cc.major)+'.'+str(cc.minor) import os import tensorrt as trt import numpy as np import time from loaders.video_loaders import decoder import torch from torchvision.transforms import v2 from transforms.model_transforms import det_transforms, clip_transforms from loaders.model_loaders import TensorRTModel from common_code import file_names import logging log = logging.getLogger() clip_frame_skip_interval = 8 clip_cropped_frame_skip_interval = 24 det_frame_skip_interval = 2 det_threshold = 0.5 clip_engine_path = os.path.join( rt_path, "models/ViT-SO400M-16-SigLIP2-512_visual_fp16_"+cc_str+".engine" ) obj_engine_path = os.path.join(rt_path, "models/dfine_large_"+cc_str+".engine") species_list_file = os.path.join(rt_path, "models/dfine_large.species_keep") abspath_species_list_file = os.path.abspath(species_list_file) with open(species_list_file, "r") as ff: species_list = ff.read().split("\n") det_model = TensorRTModel(obj_engine_path) det_model.prepare() det_model.allocate_outputs() clip_model = TensorRTModel(clip_engine_path) clip_model.prepare() clip_model.allocate_outputs() def score_video_cached(file_path, return_dict = False): file_path = file_names.resolve_file_location(file_path) det_path = file_names.get_det_npz_path(file_path) emb_path = file_names.get_embed_npz_path(file_path) vc = [det_path, emb_path] do_score = True if all([os.path.exists(x) for x in vc]): do_score = False if return_dict: det_dict = dict(np.load(det_path, allow_pickle = True)) emb_dict = dict(np.load(emb_path, allow_pickle = True)) det_dict['src_path'] = det_dict['src_path'].item() det_dict['transform_stats'] = det_dict['transform_stats'].item() emb_dict['src_path'] = emb_dict['src_path'].item() n_unq_hash = len(torch.tensor(emb_dict['embeds']).hash_tensor(dim=1).unique()) n_total_vec = emb_dict['embeds'].shape[0] if n_total_vec / n_unq_hash > 1.5: log.error(f'Recreating for {emb_path} because of cyclic values') do_score = True if not do_score and return_dict: return det_dict, emb_dict if not do_score and not return_dict: return return score_video(file_path) def score_video_sub_clip_cached(enc_file_path): crop_embeds_path = file_names.get_cropped_embed_npz_path(enc_file_path) if os.path.exists(crop_embeds_path): return else: score_video_sub_clip(enc_file_path) def score_video_sub_clip(enc_file_path): crop_embeds_path = file_names.get_cropped_embed_npz_path(enc_file_path) enc_file_path = file_names.resolve_file_location(enc_file_path) nvc_batch_size = 32 decoder.reconfigure_decoder(enc_file_path) all_sc = list() keep_reading_video = True clip_src_tensor_list = [] det_src_tensor_list = [] c_frm_num = 0 st = time.time() n_det_scores = 0 n_clip_scores = 0 transform_stats = None thresh_frames = list() thresh_scores = list() thresh_labels = list() thresh_boxes = list() det_all_scored_frames = list() clip_embeddings = list() clip_frames = list() while keep_reading_video: frames = decoder.get_batch_frames(nvc_batch_size) if len(frames) == 0: keep_reading_video = False for frame in frames: if (c_frm_num % clip_cropped_frame_skip_interval) == 0: clip_src_tensor_list.append( {"tensor": torch.from_dlpack(frame), "frame_number": c_frm_num} ) c_frm_num += 1 while len(clip_src_tensor_list) >= clip_model.batch_size: clip_tens_pass = clip_src_tensor_list[0 : clip_model.batch_size] clip_src_tensor_list = clip_src_tensor_list[clip_model.batch_size :] clip_frame_numbers = [x["frame_number"] for x in clip_tens_pass] clip_frames.append(clip_frame_numbers) clip_stacked = ( torch.stack([x["tensor"] for x in clip_tens_pass], dim=0) / 255.0 ) movie_size = list(clip_stacked.shape[::-1][0:2]) crop_div = 4 olap_factor = 0.5 crop_width = int(movie_size[0]/crop_div) crop_height = int(crop_width) spacing_width = int(crop_width * (1-olap_factor)) spacing_height = int(crop_height * (1-olap_factor)) starts_width = list(range(0,movie_size[0], spacing_width)) starts_height = list(range(0, movie_size[1], spacing_height)) crop_areas = set() for st_w in starts_width: for st_h in starts_height: en_w = st_w + crop_width en_h = st_h + crop_height if en_w > movie_size[0]: st_w = movie_size[0] - crop_width en_w = st_w + crop_width if en_h > movie_size[1]: st_h = movie_size[1] - crop_height en_h = st_h + crop_height crop_area = ( st_w, st_h, en_w, en_h) crop_areas.add(crop_area) cropped_embeddings = dict() for crop_area in crop_areas: ( st_w, st_h, en_w, en_h) = crop_area clip_stack_score = clip_stacked[:,:, st_h:en_h, st_w:en_w] clip_res = clip_model.score(clip_transforms(clip_stack_score).data_ptr()) cropped_embeddings[crop_area] = np.copy(clip_res["projected"]) n_clip_scores += clip_res["projected"].shape[0] clip_embeddings.append(cropped_embeddings) clip_dict = dict() clip_dict["src_path"] = enc_file_path cropped_embeddings = dict() for c_clip in clip_embeddings: for k,v in c_clip.items(): if k not in cropped_embeddings: cropped_embeddings[k] = list() cropped_embeddings[k].append(v) clip_dict["frame_numbers"] =np.concatenate(clip_frames) clip_dict["embeds_cropped"] = {k:np.concatenate(v) for k,v in cropped_embeddings.items()} np.savez(crop_embeds_path, **clip_dict) def score_video(enc_file_path): enc_file_path = file_names.resolve_file_location(enc_file_path) nvc_batch_size = 32 decoder.reconfigure_decoder(enc_file_path) all_sc = list() keep_reading_video = True clip_src_tensor_list = [] det_src_tensor_list = [] c_frm_num = 0 st = time.time() n_det_scores = 0 n_clip_scores = 0 transform_stats = None thresh_frames = list() thresh_scores = list() thresh_labels = list() thresh_boxes = list() det_all_scored_frames = list() clip_embeddings = list() clip_frames = list() while keep_reading_video: frames = decoder.get_batch_frames(nvc_batch_size) if len(frames) == 0: keep_reading_video = False for frame in frames: if (c_frm_num % clip_frame_skip_interval) == 0: clip_src_tensor_list.append( {"tensor": torch.from_dlpack(frame), "frame_number": c_frm_num} ) if (c_frm_num % det_frame_skip_interval) == 0: det_src_tensor_list.append( {"tensor": torch.from_dlpack(frame), "frame_number": c_frm_num} ) c_frm_num += 1 while len(det_src_tensor_list) >= det_model.batch_size: det_tens_pass = det_src_tensor_list[0 : det_model.batch_size] det_src_tensor_list = det_src_tensor_list[det_model.batch_size :] det_frame_numbers = [x["frame_number"] for x in det_tens_pass] det_stacked = ( torch.stack([x["tensor"] for x in det_tens_pass], dim=0) / 255.0 ) det_res = det_model.score(det_transforms(det_stacked).data_ptr()) n_det_scores += det_res["scores"].shape[0] if transform_stats is None: transform_stats = det_transforms.transforms[0].transform_stats det_scores = det_res["scores"] det_labels = det_res["labels"] det_boxes = det_res["boxes"] det_all_scored_frames.append(det_frame_numbers) (n_batch, n_queries, n_regressed) = det_boxes.shape frames = np.repeat( np.asarray(det_frame_numbers)[:, None], n_queries, axis=1 ) flattened_boxes = det_boxes.reshape((n_batch * n_queries, n_regressed)) mask_to_keep = det_scores.ravel() > det_threshold thresh_frames.append(frames.ravel()[mask_to_keep]) thresh_scores.append(det_scores.ravel()[mask_to_keep]) thresh_labels.append(det_labels.ravel()[mask_to_keep]) thresh_boxes.append(flattened_boxes[mask_to_keep, :]) while len(clip_src_tensor_list) >= clip_model.batch_size: clip_tens_pass = clip_src_tensor_list[0 : clip_model.batch_size] clip_src_tensor_list = clip_src_tensor_list[clip_model.batch_size :] clip_frame_numbers = [x["frame_number"] for x in clip_tens_pass] clip_frames.append(clip_frame_numbers) clip_stacked = ( torch.stack([x["tensor"] for x in clip_tens_pass], dim=0) / 255.0 ) clip_res = clip_model.score(clip_transforms(clip_stacked).data_ptr()) clip_embeddings.append(np.copy(clip_res["projected"])) n_clip_scores += clip_res["projected"].shape[0] det_dict = dict() det_dict["src_path"] = enc_file_path det_dict["transform_stats"] = transform_stats det_dict["scored_frames"] = np.concatenate(det_all_scored_frames) det_dict["final_frames"] = np.concatenate(thresh_frames) det_dict["final_labels"] = np.concatenate(thresh_labels) det_dict["species_label_map"] = { int(x): species_list[x] for x in np.unique(det_dict["final_labels"]) } det_dict["final_scores"] = np.concatenate(thresh_scores) det_dict["final_boxes"] = np.concatenate(thresh_boxes) clip_dict = dict() clip_dict["src_path"] = enc_file_path clip_dict["embeds"] = np.concatenate(clip_embeddings).astype(np.float16) clip_dict["frame_numbers"] = np.concatenate(clip_frames) np.savez(file_names.get_embed_npz_path(enc_file_path), **clip_dict) np.savez(file_names.get_det_npz_path(enc_file_path), **det_dict) return det_dict, clip_dict