from PIL import Image, ImageDraw, ImageFont from itertools import accumulate import numpy as np import torch import copy import imageio from einops import rearrange import cv2 from sklearn.decomposition import PCA from flow_viz import flow_to_rgb from fourm.utils.plotting_utils import _get_optimal_text_placement import albumentations as A import os import torchvision.transforms.functional as TF import einops import torch.nn.functional as F def expand_to_aspect_ratio(input_shape, target_aspect_ratio=None): """Increase the size of the bounding box to match the target shape.""" if target_aspect_ratio is None: return input_shape try: w, h = input_shape except (ValueError, TypeError): return input_shape w_t, h_t = target_aspect_ratio if h / w < h_t / w_t: h_new = max(w * h_t / w_t, h) w_new = w else: h_new = h w_new = max(h * w_t / h_t, w) if h_new < h or w_new < w: breakpoint() return np.array([w_new, h_new]) def fill_masked_positions_with_random(tensor, mask, min_val=0, max_val=32000): """ Fill positions in a tensor with random values where the mask is True. Args: tensor (torch.Tensor): The input tensor to modify mask (torch.Tensor): Boolean mask of same shape as tensor min_val (int): Minimum random value (inclusive) max_val (int): Maximum random value (exclusive) Returns: torch.Tensor: The modified tensor """ # Make a copy of the original tensor to avoid modifying the input result = tensor.clone() # Generate random values for the masked positions # We convert mask to int64 to make random_values have the right shape num_masked = mask.sum().item() if num_masked > 0: random_values = torch.randint(min_val, max_val, (num_masked,), dtype=tensor.dtype, device=tensor.device) # Apply random values where mask is True result[mask] = random_values return result def check_video_presence(base_dir, base_video_name): # Walk through all subfolders recursively for root, _, files in os.walk(base_dir): if base_video_name in files: print(f"Found in: {os.path.join(root, base_video_name)}") # optional debug return True return False def save_video_pairs_as_image(sorted_video_tensors_gt, sorted_video_tensors_pred, output_path='video_comparison.png', gap_size=10, separator_after_frame=None, separator_width=3): """ Save video tensor pairs as a single image with frames arranged horizontally. Args: sorted_video_tensors_gt: List of ground truth video tensors, each with shape (17, H, W, C) sorted_video_tensors_pred: List of predicted video tensors, each with shape (17, H, W, C) output_path: Path to save the output image gap_size: Size of the gap (in pixels) between video pairs separator_after_frame: If set, insert a vertical green line after this frame index to mark the boundary between conditioning frames and predicted frames. separator_width: Width in pixels of the green separator line (default: 3) """ pair_rows = [] def _make_row(video, T, h, c): frames = [np.array(video[i]) for i in range(T)] if separator_after_frame is not None and 0 < separator_after_frame < T: sep = np.zeros((h, separator_width, c), dtype=np.uint8) sep[:, :, 1] = 255 # green return np.concatenate(frames[:separator_after_frame] + [sep] + frames[separator_after_frame:], axis=1) return np.concatenate(frames, axis=1) # Process each video pair for idx, (gt_video, pred_video) in enumerate(zip(sorted_video_tensors_gt, sorted_video_tensors_pred)): # gt_video and pred_video have shape (17, H, W, C) T = gt_video.shape[0] h = np.array(gt_video[0]).shape[0] c = np.array(gt_video[0]).shape[2] gt_row = _make_row(gt_video, T, h, c) pred_row = _make_row(pred_video, T, h, c) # Stack GT and Pred vertically for this pair pair = np.concatenate([gt_row, pred_row], axis=0) # Concatenate along height pair_rows.append(pair) # Add gap after each pair except the last one if idx < len(sorted_video_tensors_gt) - 1: # Create a white gap with same width as the pair gap = np.ones((gap_size, pair.shape[1], pair.shape[2]), dtype=pair.dtype) * 255 pair_rows.append(gap) # Concatenate all pairs vertically final_image = np.concatenate(pair_rows, axis=0) # Ensure values are in valid range [0, 255] for uint8 if final_image.dtype == np.float32 or final_image.dtype == np.float64: if final_image.max() <= 1.0: final_image = (final_image * 255).astype(np.uint8) else: final_image = np.clip(final_image, 0, 255).astype(np.uint8) # Convert to PIL Image and save print(final_image.shape) final_image = np.clip(final_image, 0, 255).astype(np.uint8) # <-- Add this line img = Image.fromarray(final_image) img.save(output_path) print(f"Saved visualization to {output_path}") print(f"Final image shape: {final_image.shape}") return final_image def concat_videos_horizontally(*videos, gap=10): """ Concatenates multiple video tensors horizontally with a gap. Args: *videos: Sequence of video tensors of shape (T, H, W, C) or None. Videos must be provided left-to-right. Once a None is encountered, all subsequent videos must also be None. gap (int): Width of the gap in pixels (default: 10). Returns: Tensor of shape (T, H, W_new, C) with videos concatenated horizontally. """ # Filter videos until the first None, then stop valid_videos = [] for v in videos: if v is None: break valid_videos.append(v) # Ensure no "gaps" like [video0, video1, None, video3] if any(v is not None for v in videos[len(valid_videos):]): raise ValueError("Videos must be filled left-to-right without gaps (no None in between).") # Check shape consistency T, H, W, C = valid_videos[0].shape for v in valid_videos: if v.shape[0] != T or v.shape[1] != H or v.shape[3] != C: raise ValueError("All videos must have the same (T, H, C). Width can differ.") # Build concatenation with gaps device = valid_videos[0].device dtype = valid_videos[0].dtype gap_tensor = torch.zeros((T, H, gap, C), dtype=dtype, device=device) combined = [valid_videos[0]] for v in valid_videos[1:]: combined.append(gap_tensor) combined.append(v) return torch.cat(combined, dim=2) # concatenate along width def concat_videos_horizontally_np(*videos, gap=10): """ Concatenates multiple video arrays horizontally with a gap. Args: *videos: Sequence of numpy arrays of shape (T, H, W, C) or None. Videos must be provided left-to-right. Once a None is encountered, all subsequent entries must also be None. gap (int): Width of the gap in pixels. Returns: numpy array of shape (T, H, W_new, C) """ # Filter videos until the first None, then stop valid_videos = [] for v in videos: if v is None: break valid_videos.append(v) # Ensure no gaps after a None if any(v is not None for v in videos[len(valid_videos):]): raise ValueError("Videos must be given left-to-right without gaps (no None in between).") # Shape consistency check T, H, W, C = valid_videos[0].shape for v in valid_videos: if v.shape[0] != T or v.shape[1] != H or v.shape[3] != C: raise ValueError("All videos must have the same (T, H, C). Width can differ.") # Gap array gap_array = np.zeros((T, H, gap, C), dtype=valid_videos[0].dtype) # Construct final list combined = [valid_videos[0]] for v in valid_videos[1:]: combined.append(gap_array) combined.append(v) return np.concatenate(combined, axis=2) # concatenate along width def scale_bbox_dict_with_crop(data_dict, orig_size, new_size, crop_coords=None, bbox_shape=[192, 256]): """ Scale bbox dictionary, optionally applying crop like in humanposes_crop_and_resize. Args: data_dict: dict with "bbox_xyxy" etc. orig_size: (W, H) of original new_size: (W, H) of target crop_coords: (top, left, h, w) if cropping applied, else None bbox_shape: target aspect ratio for bbox expansion, default [192, 256] """ if data_dict is None: return None scaled_dict = copy.deepcopy(data_dict) W_orig, H_orig = orig_size W_new, H_new = new_size # Normalize bboxes = scaled_dict["bbox_xyxy"].clone() bboxes[:, [0, 2]] /= W_orig bboxes[:, [1, 3]] /= H_orig # Apply crop in normalized coords if crop_coords is not None: top, left, h, w = crop_coords xmin, ymin, xmax, ymax = left / W_orig, top / H_orig, (left + w) / W_orig, (top + h) / H_orig bboxes[:, [0, 2]] = (bboxes[:, [0, 2]] - xmin) / (xmax - xmin) bboxes[:, [1, 3]] = (bboxes[:, [1, 3]] - ymin) / (ymax - ymin) # Clip # bboxes = torch.clamp(bboxes, 0, 1) # Map back to target pixel space bboxes[:, [0, 2]] *= W_new bboxes[:, [1, 3]] *= H_new scaled_dict["bbox_xyxy"] = bboxes # Recompute center cx = (bboxes[:, 0] + bboxes[:, 2]) / 2 cy = (bboxes[:, 1] + bboxes[:, 3]) / 2 scaled_dict["box_center"] = torch.stack([cx, cy], dim=1) # Compute box_size following the original logic box_sizes = [] for i in range(bboxes.shape[0]): box_w = (bboxes[i, 2] - bboxes[i, 0]).item() box_h = (bboxes[i, 3] - bboxes[i, 1]).item() # Follow original logic: scale = bbox_dims / 200, then bbox_size = expand(scale*200).max() scale = np.array([box_w, box_h]) / 200.0 expanded_dims = expand_to_aspect_ratio(scale * 200.0, target_aspect_ratio=bbox_shape) bbox_size = expanded_dims.max() box_sizes.append(bbox_size) scaled_dict["box_size"] = torch.tensor(box_sizes, dtype=torch.float32) # Update img_size B = bboxes.shape[0] scaled_dict["img_size"] = torch.tensor([W_new, H_new], dtype=torch.float32).repeat(B, 1) return scaled_dict def save_video(frames, output_video_path, fps=4): with imageio.get_writer(output_video_path, fps=fps, format=".mp4", quality=10, macro_block_size=None) as writer: for frame in frames: writer.append_data(frame) def to_cuda_recursive(d): """Recursively move all tensors in a nested dict to CUDA.""" for k, v in d.items(): if isinstance(v, torch.Tensor): d[k] = v.cuda().float() elif isinstance(v, dict): d[k] = to_cuda_recursive(v) return d def filter_dict_by_crop(data_dict, crop_coords, orig_width, orig_height): """ Filter dictionary entries based on crop coordinates. Args: data_dict: Dictionary with batched tensors crop_coords: Tuple of (top, left, height, width) for crop region orig_width: Original image width orig_height: Original image height Returns: Filtered dictionary with only valid entries """ top, left, h, w = crop_coords bbox_xyxy = data_dict['bbox_xyxy'] # Convert to numpy if it's a tensor if isinstance(bbox_xyxy, torch.Tensor): bbox_xyxy_np = bbox_xyxy.cpu().numpy() else: bbox_xyxy_np = bbox_xyxy.copy() valid_indices = [] # Check each bbox for i in range(len(bbox_xyxy_np)): bbox_curr = bbox_xyxy_np[i].copy() # Normalize bbox coordinates bbox_curr[0::2] = bbox_curr[0::2] / orig_width # x coordinates bbox_curr[1::2] = bbox_curr[1::2] / orig_height # y coordinates # Define crop region in normalized coordinates xmin, ymin, xmax, ymax = left, top, left + w, top + h bbox_curr = A.bbox_crop(bbox_curr, x_min=xmin, y_min=ymin, x_max=xmax, y_max=ymax, rows=orig_height, cols=orig_width) bbox_curr = np.array(bbox_curr) # Check if bbox is out of range if (np.all(bbox_curr[1::2] < 0) or np.all(bbox_curr[0::2] < 0) or np.all(bbox_curr[1::2] > 1.0) or np.all(bbox_curr[0::2] > 1.0)): print("Skipping out-of-range bbox") continue valid_indices.append(i) # Filter all dictionary entries using valid indices if not valid_indices: print("No valid bboxes found after filtering") return None filtered_dict = {} for key, value in data_dict.items(): if key == 'pred_smpl_params': filtered_dict[key] = {} for new_key, new_value in value.items(): filtered_dict[key][new_key] = new_value[valid_indices] elif isinstance(value, (torch.Tensor, np.ndarray)): filtered_dict[key] = value[valid_indices] elif isinstance(value, list): filtered_dict[key] = [value[i] for i in valid_indices] else: filtered_dict[key] = value return filtered_dict def convert_dict_to_tensors(all_out, device='cuda'): """ Convert a dictionary with list values back to tensors. Args: all_out (dict): Dictionary with list values to convert device (str): Device to put tensors on ('cuda', 'cpu', etc.) Returns: dict: Dictionary with tensor values """ converted_dict = {} for k, v in all_out.items(): if k == 'pred_smpl_params': # Handle nested dictionary for pred_smpl_params converted_dict[k] = {} for k2, v2 in v.items(): converted_dict[k][k2] = torch.tensor(v2, device=device) else: # Convert regular list to tensor converted_dict[k] = torch.tensor(v, device=device) return converted_dict def vertical_concat_videos(video_arrays): """ Vertically concatenate a list of video arrays. Args: video_arrays (list): List of numpy arrays with shape [17, H, W, C] Returns: np.ndarray: A single array with videos stacked vertically Shape will be [17, sum(H), W, C] """ # Check if list is empty if not video_arrays: return None # Ensure all videos have the same number of frames, width and channels frames, heights, widths, channels = zip(*[v.shape for v in video_arrays]) # Verify all videos have the same number of frames, width and channels if len(set(frames)) > 1 or len(set(widths)) > 1 or len(set(channels)) > 1: raise ValueError("All videos must have the same number of frames, width, and channels") # Stack each frame vertically result = [] for f in range(frames[0]): frame_list = [video[f] for video in video_arrays] stacked_frame = np.concatenate(frame_list, axis=0) # Concat along height dimension result.append(stacked_frame) # Stack all frames back together return np.stack(result) def create_frame_ids(x): # Define special token range special_min, special_max = 30004, 30020 # Initialize frame ID tensor frame_ids = torch.zeros_like(x) # We'll process each batch element for b in range(x.size(0)): tokens = x[b] special_mask = (tokens >= special_min) & (tokens <= special_max) special_indices = torch.nonzero(special_mask, as_tuple=False).squeeze(-1) current_id = 0 for i, idx in enumerate(special_indices): start = idx end = special_indices[i + 1] if i + 1 < len(special_indices) else len(tokens) frame_ids[b, start:end] = current_id current_id += 1 return frame_ids def image_mask_first_frame_conditional(tensor: torch.Tensor, GT_tokens: int, input_budget: int, target_budget: int): """Applies input and target masking to an image tensor sequentially Args: tensor: Image tensor GT_tokens: Number of tokens in the tensor input_budget: Token budget for the input target_budget: Token budget for the target Returns: Dictionary containing the masked image tensor, the input mask, the target mask, and the decoder attention mask """ # Input mask: First `input_budget` tokens are not masked (0), rest are masked (1) input_mask = torch.ones(GT_tokens, dtype=torch.bool) input_mask[:input_budget] = 0 # First `input_budget` positions are not masked # Target mask: The next `target_budget` tokens are not masked (0), rest are masked (1) target_mask = torch.ones(GT_tokens, dtype=torch.bool) if target_budget is not None: target_mask[input_budget:input_budget + target_budget] = 0 # Next `target_budget` positions are not masked else: target_mask = ~input_mask # If target_budget is None, complement input_mask # Compute decoder attention mask decoder_attention_mask = torch.zeros(GT_tokens, dtype=torch.int) first_mask_token = torch.argmin(target_mask + torch.arange(target_mask.shape[0], device=target_mask.device) * 1e-6) decoder_attention_mask[first_mask_token] = (~target_mask).sum() # Equivalent to target budget return { "tensor": torch.tensor(tensor).long().cuda(), "input_mask": input_mask.unsqueeze(0).cuda(), "target_mask": target_mask.unsqueeze(0).cuda(), "decoder_attention_mask": decoder_attention_mask.unsqueeze(0).cuda(), } def transform_tensor_with_markers(tensor, start_sentinel=5, end_sentinel=21, start_marker=30004, num_markers=17): """ Transform tensor by replacing sentinel tokens with special markers and extending to always have num_markers frame markers. Args: tensor: Input tensor containing sentinel tokens start_sentinel: First sentinel token value (default: 5) end_sentinel: Last sentinel token value (default: 21) start_marker: First special marker value (default: 30004) num_markers: Total number of frame markers to ensure (default: 17) """ device = tensor.device result = [] # Find all sentinel tokens in the tensor sentinel_positions = {} flat_tensor = tensor.flatten() for i, token in enumerate(flat_tensor): token_val = token.item() if start_sentinel <= token_val <= end_sentinel: if token_val not in sentinel_positions: sentinel_positions[token_val] = [] sentinel_positions[token_val].append(i) # Get sorted sentinel tokens that exist in the tensor existing_sentinels = sorted(sentinel_positions.keys()) # Process existing sentinels last_pos = 0 for sentinel_idx, sentinel_val in enumerate(existing_sentinels): # Get the position of this sentinel sentinel_pos = sentinel_positions[sentinel_val][0] # Add the marker for this sentinel marker = start_marker + sentinel_idx result.append(marker) # Find tokens between this sentinel and the next (or end) if sentinel_idx < len(existing_sentinels) - 1: next_sentinel_pos = sentinel_positions[existing_sentinels[sentinel_idx + 1]][0] else: next_sentinel_pos = len(flat_tensor) # Add tokens between sentinels (excluding the sentinel token itself) for j in range(sentinel_pos + 1, next_sentinel_pos): if flat_tensor[j].item() not in range(start_sentinel, end_sentinel + 1): result.append(flat_tensor[j].item()) # Complete with remaining markers if we have fewer than num_markers num_existing = len(existing_sentinels) if num_existing < num_markers: for i in range(num_existing, num_markers): marker = start_marker + i # Reset sentinel numbering to start from start_sentinel sentinel = start_sentinel + (i - num_existing) result.append(marker) result.append(sentinel) return torch.tensor(result, device=device).unsqueeze(0) def pop_conditioning_domain(cond_domains, target_domains, partial_conditioning_tokens, complete_partial_conditioned_modalities, *other_lists): for cond_domain, partial_tokens in zip(cond_domains, partial_conditioning_tokens): if cond_domain not in target_domains: raise ValueError(f"{cond_domain} not found in target_domains") idx = target_domains.index(cond_domain) popped_values = [] # Remove cond_domain from target_domains only if partial tokens are none, means full, and user wants to complete it if (partial_tokens is None) or not complete_partial_conditioned_modalities: target_domains.pop(idx) # Remove corresponding elements from each of the other lists popped_values = [lst.pop(idx) for lst in other_lists] return target_domains, *other_lists, popped_values def create_text_frame(text, height, width=50, font_size=45, bg_color=(0, 0, 0), text_color=(255, 255, 255)): """ Create a static frame with text. Args: text (str): Text to display height (int): Height of the frame (should match video height) width (int): Width of the text frame font_size (int): Font size for the text bg_color (tuple): Background color (R, G, B) text_color (tuple): Text color (R, G, B) Returns: np.ndarray: Text frame with shape [height, width, 3] """ # Create PIL image img = Image.new('RGB', (width, height), bg_color) draw = ImageDraw.Draw(img) try: # Try to use a better font if available font = ImageFont.truetype("arial.ttf", font_size) except: # Fall back to default font font = ImageFont.load_default() # Get text bounding box for centering bbox = draw.textbbox((0, 0), text, font=font) text_width = bbox[2] - bbox[0] text_height = bbox[3] - bbox[1] # Calculate position to center the text x = (width - text_width) // 2 y = (height - text_height) // 2 # Draw text draw.text((x, y), text, fill=text_color, font=font) # Convert PIL image to numpy array return np.array(img) def resize_and_duplicate(arr, target_size=(128, 128)): """ arr: numpy array of shape (T, H, W, C) """ T, H, W, C = arr.shape # ---- 1. Resize ---- resized = np.zeros((T, target_size[0], target_size[1], C), dtype=arr.dtype) for i in range(T): resized[i] = cv2.resize(arr[i], target_size, interpolation=cv2.INTER_NEAREST) # ---- 2. Duplicate if T == 8 ---- if T == 8: idx = [0, 0, 0] + sum([[i, i] for i in range(1, T)], []) resized = resized[idx] return resized def resize_and_duplicate_batched(arr, target_size=(128, 128)): """ arr: numpy array of shape (B, T, H, W, C) or (T, H, W, C) Returns: numpy array of same batch structure with resized frames """ # Handle both batched and unbatched inputs if arr.ndim == 4: arr = arr[np.newaxis, ...] # Add batch dim unbatch_output = True else: unbatch_output = False B, T, H, W, C = arr.shape # ---- 1. Resize using PyTorch (vectorized) ---- # Convert to torch: (B, T, H, W, C) -> (B*T, C, H, W) arr_torch = torch.from_numpy(arr).reshape(B * T, H, W, C).permute(0, 3, 1, 2) # Resize all frames at once resized_torch = F.interpolate( arr_torch.float(), size=target_size, mode='nearest' ) # Back to numpy: (B*T, C, H, W) -> (B, T, target_H, target_W, C) resized = resized_torch.permute(0, 2, 3, 1).reshape( B, T, target_size[0], target_size[1], C ).numpy().astype(arr.dtype) # ---- 2. Duplicate if T == 8 (vectorized with indexing) ---- if T == 8: idx = [0, 0, 0] + sum([[i, i] for i in range(1, T)], []) resized = resized[:, idx] # Works for entire batch # Remove batch dim if input was unbatched if unbatch_output: resized = resized[0] return resized class FeatureToPCAConverter: """Convert high-dimensional features to RGB using PCA""" def __init__(self): self.pca = None self.fitted = False def convert_to_rgb(self, features): """ Convert features [B, C, T, H, W] to RGB [B, T, H, W, 3] using PCA """ with torch.no_grad(): B, C, T, H, W = features.shape device = features.device # Move to CPU for PCA feats_reshaped = features.permute(0, 2, 3, 4, 1).reshape(-1, C).cpu().numpy() # Fit PCA on first call if not self.fitted: self.pca = PCA(n_components=3) self.pca.fit(feats_reshaped) self.fitted = True print(f"PCA fitted - Variance explained: {self.pca.explained_variance_ratio_.sum():.3f}") # PCA transform pca_features = self.pca.transform(feats_reshaped) # Back to tensor rgb = torch.from_numpy(pca_features).float().reshape(B, T, H, W, 3) rgb = rgb.to(device) # Normalize per batch to [0, 1] flat = rgb.view(B, -1, 3) rgb_min = flat.min(dim=1)[0].view(B, 1, 1, 1, 3) rgb_max = flat.max(dim=1)[0].view(B, 1, 1, 1, 3) rgb = (rgb - rgb_min) / (rgb_max - rgb_min + 1e-8) return rgb def convert_dict_to_tensors(all_out, device='cpu'): """ Convert a dictionary with list values back to tensors. Args: all_out (dict): Dictionary with list values to convert device (str): Device to put tensors on ('cuda', 'cpu', etc.) Returns: dict: Dictionary with tensor values """ converted_dict = {} for k, v in all_out.items(): if k == 'pred_smpl_params': # Handle nested dictionary for pred_smpl_params converted_dict[k] = {} for k2, v2 in v.items(): converted_dict[k][k2] = torch.tensor(v2, device=device) else: # Convert regular list to tensor converted_dict[k] = torch.tensor(v, device=device) return converted_dict def load_and_decode_tokens_batched(tokens, tokenizer_model, device='cuda', original_len=17, crop_settings=None, modality="DINOv2-B14"): """ Load saved tokens and decode them to reconstruction Args: tokens: Token indices [num_crops, num_tokens] or [num_tokens] tokenizer_model: Loaded VidTok model device: Device to use original_len: Original temporal length for decoding Returns: Reconstructed video [B, C, T, H, W] """ # Convert to tensor if numpy if isinstance(tokens, np.ndarray): tokens = torch.from_numpy(tokens) tokens = tokens.to(device) if modality == "vjepa": # if original_len == 8: t = 4 # TODO: hard coded for now padding_applied = 0 else: t = 5 padding_applied = 3 h = w = int((tokens.shape[1] / t) ** 0.5) tokens = rearrange(tokens, "b (t h w) ->b t h w", t=t, h=h, w=w) with torch.no_grad(), torch.autocast(device_type='cuda', dtype=torch.float16): # Decode from indices reconstructed = tokenizer_model.decode(tokens, decode_from_indices=True, original_len=original_len, padding_applied=padding_applied) return reconstructed.cpu() def load_and_decode_tokens(tokens, tokenizer_model, device='cuda', original_len=17, crop_settings=None, modality="DINOv2-B14"): """ Load saved tokens and decode them to reconstruction Args: tokens: Token indices [num_crops, num_tokens] or [num_tokens] tokenizer_model: Loaded VidTok model device: Device to use original_len: Original temporal length for decoding Returns: Reconstructed video [B, C, T, H, W] """ # Convert to tensor if numpy if isinstance(tokens, np.ndarray): tokens = torch.from_numpy(tokens) tokens = tokens.to(device) if modality == "vjepa": # if original_len == 8: t = 4 # TODO: hard coded for now padding_applied = 0 else: t = 5 padding_applied = 3 h = w = int((tokens.shape[0] / t) ** 0.5) tokens = rearrange(tokens, "(t h w) -> t h w", t=t, h=h, w=w) tokens = tokens.unsqueeze(0) with torch.no_grad(), torch.autocast(device_type='cuda', dtype=torch.float16): # Decode from indices reconstructed = tokenizer_model.decode(tokens, decode_from_indices=True, original_len=original_len, padding_applied=padding_applied) return reconstructed.cpu() def get_class_name_by_index(file_path, index): """ Get class name by index from a text file. Args: file_path (str): Path to the text file containing class names index (int): Index of the class name to retrieve (0-based) Returns: str: Class name at the given index, or None if index is out of range """ try: with open(file_path, 'r', encoding='utf-8') as file: lines = file.readlines() # Strip whitespace from each line class_names = [line.strip() for line in lines if line.strip()] # Check if index is valid if 0 <= index < len(class_names): return class_names[index] else: print(f"Index {index} is out of range. Valid range: 0-{len(class_names) - 1}") return None except FileNotFoundError: print(f"File not found: {file_path}") return None except Exception as e: print(f"Error reading file: {e}") return None def load_class_names(file_path): """ Load all class names into a list for multiple lookups. More efficient if you need to do many lookups. Args: file_path (str): Path to the text file containing class names Returns: list: List of class names, or empty list if error """ try: with open(file_path, 'r', encoding='utf-8') as file: return [line.strip() for line in file if line.strip()] except Exception as e: print(f"Error loading class names: {e}") return [] def vertical_concat_videos_with_labels(video_arrays, labels=None, text_width=100, font_size=60, line_thickness=5): """ Vertically concatenate a list of video arrays with text labels on the left, adding a green line separator between videos. Args: video_arrays (list): List of numpy arrays with shape [T, H, W, C] labels (list): List of text labels for each video. If None, uses default labels. text_width (int): Width of the text label area font_size (int): Font size for labels line_thickness (int): Thickness of the green separator line Returns: np.ndarray: A single array with labeled videos stacked vertically """ if not video_arrays: return None if labels is None: labels = ["Only class"] + [f"{tokens} tokens" for tokens in [256, 512]] if len(labels) < len(video_arrays): for i in range(len(labels), len(video_arrays)): labels.append(f"Video {i + 1}") # Get shape info frames, heights, widths, channels = zip(*[v.shape for v in video_arrays]) if len(set(frames)) > 1 or len(set(widths)) > 1 or len(set(channels)) > 1: raise ValueError("All videos must have the same number of frames, width, and channels") labeled_videos = [] for i, (video, label) in enumerate(zip(video_arrays, labels)): text_frame = create_text_frame( text=label, height=video.shape[1], width=text_width, font_size=font_size ) labeled_video_frames = [] for frame_idx in range(video.shape[0]): labeled_frame = np.concatenate([text_frame, video[frame_idx]], axis=1) labeled_video_frames.append(labeled_frame) labeled_video = np.stack(labeled_video_frames) labeled_videos.append(labeled_video) # Now vertically concatenate with green line separators result = [] green_color = np.array([0, 255, 0], dtype=np.uint8) # Green (BGR or RGB depending on your video) for f in range(frames[0]): frame_list = [] for i, video in enumerate(labeled_videos): frame_list.append(video[f]) if i < len(labeled_videos) - 1: # Add separator except after the last one sep = np.full( (line_thickness, video[f].shape[1], video[f].shape[2]), green_color, dtype=np.uint8 ) frame_list.append(sep) stacked_frame = np.concatenate(frame_list, axis=0) result.append(stacked_frame) return np.stack(result) def tensor_to_uint8(tensor): tensor = torch.clamp(tensor, -1.0, 1.0) tensor = (tensor + 1.0) / 2.0 # -1,1 -> 0,1; c,h,w tensor = (tensor.cpu().numpy() * 255).astype(np.uint8) return tensor def merge_tokens_with_frames_no_eos(tokens, frames, counts): """Merge tokens with frame identifiers efficiently.""" if not tokens or not frames or not counts: return "" output = [] indices = [0] + list(accumulate(counts)) for i in range(len(counts)): frame_tokens = ' '.join(tokens[indices[i]:indices[i + 1]]) merged = f"{frames[i]} {frame_tokens}" if frame_tokens else frames[i] output.append(merged) return ' '.join(output) def merge_detection_tokens_with_sentinel_tokens_no_eos(tokens, sentinel_tokens, counts): """Merge tokens with only sentinel tokens""" if not tokens or not sentinel_tokens or not counts: return "" output = [] indices = [0] + list(accumulate(counts)) for i in range(len(counts)): frame_tokens = ' '.join(tokens[indices[i]:indices[i + 1]]) merged = f"{sentinel_tokens[i]} {frame_tokens}" if frame_tokens else sentinel_tokens[i] output.append(merged) return ' '.join(output) def merge_detection_tokens_with_sentinel_tokens(tokens, sentinel_tokens, counts): """Merge tokens with only sentinel tokens""" # if not tokens or not sentinel_tokens or not counts: # return "" output = [] indices = [0] + list(accumulate(counts)) for i in range(len(counts)): frame_tokens = ' '.join(tokens[indices[i]:indices[i + 1]]) merged = f"{sentinel_tokens[i]} {frame_tokens} [EOS]" if frame_tokens else f"{sentinel_tokens[i]} [EOS]" output.append(merged) return ' '.join(output) def reorder_lists_by_conditioning_domain(conditioning_domain, target_domains, *lists): """ Reorder all lists to put the conditioning domain's corresponding values first. Args: conditioning_domain: The domain to use as conditioning (e.g., 'det') target_domains: List of domain names *lists: Variable number of lists to reorder Returns: Tuple of reordered lists in the same order as input """ if conditioning_domain not in target_domains: raise ValueError(f"Conditioning domain '{conditioning_domain}' not found in target_domains") # Find the index of the conditioning domain cond_index = target_domains.index(conditioning_domain) # Reorder target_domains first new_target_domains = [target_domains[cond_index]] + target_domains[:cond_index] + target_domains[cond_index + 1:] # Reorder all other lists reordered_lists = [] for lst in lists: if lst is None or len(lst) != len(target_domains): # Handle None or mismatched length lists reordered_lists.append(lst) else: new_list = [lst[cond_index]] + lst[:cond_index] + lst[cond_index + 1:] reordered_lists.append(new_list) return new_target_domains, *reordered_lists def reorder_lists_by_conditioning_domains(conditioning_domains, target_domains, *lists): """ Reorder all lists to put the conditioning domains' corresponding values first. Args: conditioning_domains: List of domains to use as conditioning (e.g., ['det', 'seg']) target_domains: List of all domain names *lists: Variable number of lists to reorder Returns: Tuple of reordered lists in the same order as input """ # Ensure conditioning_domains is a list if isinstance(conditioning_domains, str): conditioning_domains = [conditioning_domains] # Check that all conditioning domains exist in target_domains missing = [d for d in conditioning_domains if d not in target_domains] if missing: raise ValueError(f"Conditioning domains {missing} not found in target_domains") # Get indices for conditioning domains and non-conditioning domains cond_indices = [target_domains.index(d) for d in conditioning_domains] remaining_indices = [i for i in range(len(target_domains)) if i not in cond_indices] # Reorder target_domains new_target_domains = [target_domains[i] for i in cond_indices + remaining_indices] # Reorder all other lists reordered_lists = [] for lst in lists: if lst is None or len(lst) != len(target_domains): reordered_lists.append(lst) else: new_list = [lst[i] for i in cond_indices + remaining_indices] reordered_lists.append(new_list) return new_target_domains, *reordered_lists def merge_tokens_with_frames(tokens, frames, counts): """Merge tokens with frame identifiers efficiently.""" # if not tokens or not frames or not counts: # return "" output = [] indices = [0] + list(accumulate(counts)) for i in range(len(counts)): frame_tokens = ' '.join(tokens[indices[i]:indices[i + 1]]) merged = f"{frames[i]} {frame_tokens} [EOS]" if frame_tokens else f"{frames[i]} [EOS]" output.append(merged) return ' '.join(output) def image_mask(tensor: torch.Tensor, GT_tokens: int, input_budget: int, target_budget: int): """Applies input and target masking to an image tensor Args: tensor: Image tensor GT_tokens: Number of tokens in the tensor input_budget: Token budget for the input target_budget: Token budget for the target Returns: Dictionary containing the masked image tensor, the input mask, the target mask, and the decoder attention mask """ # Use fixed seed for deterministic ordering across different calls torch.manual_seed(42) noise = torch.rand(GT_tokens) ids_shuffle = torch.argsort(noise, dim=0) input_mask = torch.ones(GT_tokens, dtype=torch.bool) input_mask[:input_budget] = 0 input_mask = torch.gather(input_mask, dim=0, index=ids_shuffle) if target_budget is None: target_mask = ~input_mask else: target_mask = torch.ones(GT_tokens, dtype=torch.bool) target_mask[input_budget:input_budget + target_budget] = 0 target_mask = torch.gather(target_mask, dim=0, index=ids_shuffle) decoder_attention_mask = torch.zeros(GT_tokens, dtype=torch.int) first_mask_token = torch.argmin(target_mask + torch.arange(target_mask.shape[0], device=target_mask.device) * 1e-6) decoder_attention_mask[first_mask_token] = (~target_mask).sum() # Equiv. to target budget return { "tensor": torch.tensor(tensor).long().cuda(), "input_mask": input_mask.unsqueeze(0).cuda(), "target_mask": target_mask.unsqueeze(0).cuda(), "decoder_attention_mask": decoder_attention_mask.unsqueeze(0).cuda(), } def denormalize(img, mean=None, std=None): """ Denormalizes an image. Args: img (torch.Tensor): Image to denormalize. mean (tuple): Mean to use for denormalization. std (tuple): Standard deviation to use for denormalization. """ return TF.normalize( img.clone(), mean=[-m / s for m, s in zip(mean, std)], std=[1 / s for s in std] ) def denormalize_video(video, mean=None, std=None): """ Denormalizes videos. Args: video (torch.Tensor): Video to denormalize. mean (tuple): Mean to use for denormalization. std (tuple): Standard deviation to use for denormalization. """ if len(video.shape) == 4: # single video, use denormalize return denormalize(video, mean=mean, std=std) B = video.shape[0] # pack frames into the batch dimension img = einops.rearrange(video, "b c t h w -> (b t) c h w") # denormalize each frame img_norm = denormalize(img, mean=mean, std=std) # unpack videos norm_video = einops.rearrange(img_norm, "(b t) c h w -> b c t h w", b=B) return norm_video def image_mask_first_frame_conditional(tensor: torch.Tensor, GT_tokens: int, input_budget: int, target_budget: int, device='cuda'): """Applies input and target masking to an image tensor sequentially Args: tensor: Image tensor GT_tokens: Number of tokens in the tensor input_budget: Token budget for the input target_budget: Token budget for the target Returns: Dictionary containing the masked image tensor, the input mask, the target mask, and the decoder attention mask """ # Input mask: First `input_budget` tokens are not masked (0), rest are masked (1) input_mask = torch.ones(GT_tokens, dtype=torch.bool) input_mask[:input_budget] = 0 # First `input_budget` positions are not masked # Target mask: The next `target_budget` tokens are not masked (0), rest are masked (1) target_mask = torch.ones(GT_tokens, dtype=torch.bool) if target_budget is not None: target_mask[input_budget:input_budget + target_budget] = 0 # Next `target_budget` positions are not masked else: target_mask = ~input_mask # If target_budget is None, complement input_mask # Compute decoder attention mask decoder_attention_mask = torch.zeros(GT_tokens, dtype=torch.int) first_mask_token = torch.argmin(target_mask + torch.arange(target_mask.shape[0], device=target_mask.device) * 1e-6) decoder_attention_mask[first_mask_token] = (~target_mask).sum() # Equivalent to target budget # FIXED: Only include the conditioning tokens (first input_budget tokens) in the tensor # The rest should be zeros (will be predicted during generation) full_tensor = torch.zeros(GT_tokens, dtype=torch.long, device=device) full_tensor[:input_budget] = torch.tensor(tensor[:input_budget]).long().to(device) return { "tensor": full_tensor.unsqueeze(0), "input_mask": input_mask.unsqueeze(0).to(device), "target_mask": target_mask.unsqueeze(0).to(device), "decoder_attention_mask": decoder_attention_mask.unsqueeze(0).to(device), } # Function to save side-by-side video using imageio def save_video_with_imageio(frames, output_path, fps): with imageio.get_writer(output_path, fps=fps) as writer: for frame in frames: writer.append_data(frame) print(f"Side-by-side video saved at: {output_path}") def process_tsn_frame(frame_rgb: np.ndarray, bound: int) -> tuple[np.ndarray, np.ndarray]: flow_xy = tsn_rgb_to_flow(frame_rgb, bound=bound) colored = flow_to_rgb(flow_xy) return colored, flow_xy def tsn_rgb_to_flow(flow_img: np.ndarray, bound: int = 20) -> np.ndarray: """ Convert TSN-encoded optical-flow RGB frame into (H,W,2) raw flow (float32). Uses only the first two channels as x and y; third channel (if present) is ignored. """ flow_x_norm = flow_img[..., 0].astype(np.float32) flow_y_norm = flow_img[..., 1].astype(np.float32) flow_x = np.expand_dims(((flow_x_norm * (2 * bound) / 255.0) - bound), -1) flow_y = np.expand_dims(((flow_y_norm * (2 * bound) / 255.0) - bound), -1) return np.concatenate([flow_x, flow_y], axis=-1) def only_extract_indices(crop_settings): fps = 30 required_frames = [] indices_per_crop = [] for single_crop in crop_settings: start_t, end_t, i, j, h, w, h_flip = single_crop frame_indices = np.linspace( start_t * fps, end_t * fps, 17, dtype=np.int32 ) frame_indices[-1] -= 1 indices_per_crop.append(frame_indices) required_frames.extend(frame_indices) indices = sorted(set(required_frames)) return indices, indices_per_crop # default 4 FPS def convert_raw_optical_flow(frames, bound: int) -> None: """ paths keys: original, reconstructed -> Path """ all_frames_original = [] for frame_idx in (range(len(frames))): orig_frame = frames[frame_idx] # Process TSN frames to get flow visualization orig_colored, orig_raw_flow = process_tsn_frame(orig_frame, bound=bound) all_frames_original.append(orig_colored) return np.stack(all_frames_original, axis=0) def denorm_bbox(bbox, width, height): """Convert normalized bbox [x1, y1, x2, y2] → pixel coords""" x1 = int(bbox[0] * width) y1 = int(bbox[1] * height) x2 = int(bbox[2] * width) y2 = int(bbox[3] * height) return [x1, y1, x2, y2] def order_bboxes_by_dist_to_orig(detections, thresh=0.6, max_det=6): # keep only above threshold # detections = [det for det in detections if det["confidence"] >= thresh] # sort by distance of top-left corner (x1, y1) to origin # detections = sorted( # detections, # key=lambda d: d["bbox"][0]**2 + d["bbox"][1]**2 # bbox coords are normalized # ) # keep top-k return detections[:max_det] def draw_boxes_on_blank(frame_size, detections, color=(0,0,255), bg_color=(255,255,255)): """Draw detections + labels on blank image""" h, w = frame_size img = np.full((h, w, 3), bg_color, dtype=np.uint8) for det in detections: bbox = det["bbox"] x1, y1, x2, y2 = denorm_bbox(bbox, w, h) cv2.rectangle(img, (x1, y1), (x2, y2), color, 2) # --- Add label text --- label_text = f"{det['label']}" # Determine optimal text placement and size text_position, font_scale, background_rect = _get_optimal_text_placement( img, x1, y1, x2, y2, label_text ) # Draw text background rectangle if background_rect: cv2.rectangle(img, background_rect[0], background_rect[1], color, -1) # Draw the text cv2.putText( img, text=label_text, org=text_position, fontFace=cv2.FONT_HERSHEY_SIMPLEX, fontScale=font_scale, color=bg_color, lineType=cv2.LINE_AA, ) return img def crop_and_resize_detections(detections, crop, frame_size, target_size): """ Crop detections based on crop [top, left, h, w]. Normalize wrt cropped frame, then we will directly draw on resized frame. """ top, left, ch, cw = crop H, W = frame_size filtered = [] for det in detections: x1, y1, x2, y2 = denorm_bbox(det["bbox"], W, H) # Intersection with crop nx1 = max(x1, left) ny1 = max(y1, top) nx2 = min(x2, left + cw) ny2 = min(y2, top + ch) if nx1 < nx2 and ny1 < ny2: # valid intersection # Shift to cropped coords and normalize wrt cropped region adj_x1 = (nx1 - left) / cw adj_y1 = (ny1 - top) / ch adj_x2 = (nx2 - left) / cw adj_y2 = (ny2 - top) / ch new_det = { "label": det["label"], "confidence": det["confidence"], "bbox": [adj_x1, adj_y1, adj_x2, adj_y2], } filtered.append(new_det) return filtered