first version
This commit is contained in:
500
sam2/sam_extract.py
Normal file
500
sam2/sam_extract.py
Normal file
@@ -0,0 +1,500 @@
|
||||
"""
|
||||
Runs on video and tracks objects across time, generating masks and polygons for multiple frames.
|
||||
Takes one annotation of each asset and fetches the frames before and after it.
|
||||
"""
|
||||
import os
|
||||
import cv2
|
||||
import ast
|
||||
import json
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from collections import deque
|
||||
from scipy.optimize import linear_sum_assignment
|
||||
from segment_anything import sam_model_registry, SamPredictor
|
||||
|
||||
# =========================================================
|
||||
# CONFIG
|
||||
# =========================================================
|
||||
VIDEO_PATH = "/media/bs5/New Volume1/Brinda/Dataset/GPS_location/videos/2026_0330_094759_F.MP4"
|
||||
JSON_PATH = "/media/bs5/New Volume1/Brinda/Dataset/GPS_location/output/2026_0330_094759_F.json"
|
||||
OUTPUT_DIR = "/media/bs5/New Volume1/Brinda/Dataset/GPS_location/sam/"
|
||||
|
||||
FRAME_GAP = 5
|
||||
NUM_FRAMES = 5
|
||||
|
||||
SAM_CHECKPOINT = "sam_vit_h_4b8939.pth"
|
||||
MODEL_TYPE = "vit_h"
|
||||
|
||||
|
||||
# =========================================================
|
||||
# OUTPUT DIRS
|
||||
# =========================================================
|
||||
IMG_DIR = os.path.join(OUTPUT_DIR, "images")
|
||||
MASK_DIR = os.path.join(OUTPUT_DIR, "masks")
|
||||
EXTRACT_DIR = os.path.join(OUTPUT_DIR, "masked")
|
||||
os.makedirs(OUTPUT_DIR, exist_ok=True)
|
||||
os.makedirs(IMG_DIR, exist_ok=True)
|
||||
os.makedirs(MASK_DIR, exist_ok=True)
|
||||
os.makedirs(EXTRACT_DIR, exist_ok=True)
|
||||
|
||||
5
|
||||
# ================== KALMAN FILTER ==================
|
||||
|
||||
class KalmanFilter:
|
||||
"""Kalman filter for smoother track predictions"""
|
||||
def __init__(self):
|
||||
self.kf = cv2.KalmanFilter(8, 4)
|
||||
|
||||
# State transition matrix
|
||||
self.kf.transitionMatrix = np.eye(8, dtype=np.float32)
|
||||
self.kf.transitionMatrix[0, 4] = 1
|
||||
self.kf.transitionMatrix[1, 5] = 1
|
||||
self.kf.transitionMatrix[2, 6] = 1
|
||||
self.kf.transitionMatrix[3, 7] = 1
|
||||
|
||||
# Measurement matrix
|
||||
self.kf.measurementMatrix = np.zeros((4, 8), dtype=np.float32)
|
||||
self.kf.measurementMatrix[0, 0] = 1
|
||||
self.kf.measurementMatrix[1, 1] = 1
|
||||
self.kf.measurementMatrix[2, 2] = 1
|
||||
self.kf.measurementMatrix[3, 3] = 1
|
||||
|
||||
self.kf.processNoiseCov = np.eye(8, dtype=np.float32) * 0.01
|
||||
self.kf.measurementNoiseCov = np.eye(4, dtype=np.float32) * 0.5
|
||||
self.kf.errorCovPost = np.eye(8, dtype=np.float32)
|
||||
|
||||
def predict(self):
|
||||
prediction = self.kf.predict()
|
||||
return prediction[:4].flatten()
|
||||
|
||||
def update(self, measurement):
|
||||
self.kf.correct(np.array(measurement, dtype=np.float32).reshape(4, 1))
|
||||
|
||||
|
||||
# ================== TRACK CLASS ==================
|
||||
|
||||
class Track:
|
||||
def __init__(self, tlwh, score, track_id, class_id=0, class_name='unknown'):
|
||||
self.tlwh = np.asarray(tlwh, dtype=np.float32)
|
||||
self.score = score
|
||||
self.track_id = track_id
|
||||
self.class_id = class_id
|
||||
self.class_name = class_name
|
||||
self.time_since_update = 0
|
||||
self.hits = 1
|
||||
self.age = 1
|
||||
|
||||
# Kalman filter for motion prediction
|
||||
self.kf = KalmanFilter()
|
||||
self.kf.kf.statePost = np.array([*tlwh, 0, 0, 0, 0], dtype=np.float32).reshape(8, 1)
|
||||
|
||||
# Track confidence history for stability
|
||||
self.confidence_history = deque([score], maxlen=10)
|
||||
self.avg_confidence = score
|
||||
|
||||
def predict(self):
|
||||
"""Predict next position using Kalman filter"""
|
||||
self.age += 1
|
||||
self.time_since_update += 1
|
||||
predicted = self.kf.predict()
|
||||
|
||||
if self.time_since_update < 5:
|
||||
self.tlwh = predicted
|
||||
|
||||
def update(self, tlwh, score, class_id, class_name):
|
||||
"""Update track with new detection"""
|
||||
self.tlwh = np.asarray(tlwh, dtype=np.float32)
|
||||
self.score = score
|
||||
self.class_id = class_id
|
||||
self.class_name = class_name
|
||||
self.time_since_update = 0
|
||||
self.hits += 1
|
||||
|
||||
self.confidence_history.append(score)
|
||||
self.avg_confidence = np.mean(list(self.confidence_history))
|
||||
|
||||
self.kf.update(tlwh)
|
||||
|
||||
def is_stable(self):
|
||||
"""Check if track is stable enough to display"""
|
||||
return (self.hits >= 3 and
|
||||
self.avg_confidence > 0.3 and
|
||||
self.time_since_update == 0)
|
||||
|
||||
|
||||
# ================== BYTETRACK TRACKER ==================
|
||||
|
||||
class ByteTrackEnhanced:
|
||||
"""
|
||||
Enhanced ByteTrack with per-class track ID sequencing
|
||||
"""
|
||||
def __init__(self, track_thresh=0.4, match_thresh=0.5, buffer_size=60,
|
||||
high_thresh=0.5, low_thresh=0.1, class_names=None):
|
||||
self.track_thresh = track_thresh
|
||||
self.high_thresh = high_thresh
|
||||
self.low_thresh = low_thresh
|
||||
self.match_thresh = match_thresh
|
||||
self.buffer_size = buffer_size
|
||||
|
||||
self.tracks = []
|
||||
self.class_names = class_names if class_names else []
|
||||
|
||||
# Per-class track ID counters
|
||||
self.next_id_per_class = {}
|
||||
|
||||
def get_next_id(self, class_id):
|
||||
"""Get next track ID for specific class"""
|
||||
if class_id not in self.next_id_per_class:
|
||||
self.next_id_per_class[class_id] = 1
|
||||
|
||||
track_id = self.next_id_per_class[class_id]
|
||||
self.next_id_per_class[class_id] += 1
|
||||
return track_id
|
||||
|
||||
def get_class_name(self, class_id):
|
||||
"""Get class name from class ID"""
|
||||
if class_id < len(self.class_names):
|
||||
return self.class_names[class_id]
|
||||
return f'class_{class_id}'
|
||||
|
||||
def iou(self, a, b):
|
||||
"""Calculate IoU between two boxes"""
|
||||
ax, ay, aw, ah = a
|
||||
bx, by, bw, bh = b
|
||||
ax2, ay2 = ax + aw, ay + ah
|
||||
bx2, by2 = bx + bw, by + bh
|
||||
|
||||
inter_x1 = max(ax, bx)
|
||||
inter_y1 = max(ay, by)
|
||||
inter_x2 = min(ax2, bx2)
|
||||
inter_y2 = min(ay2, by2)
|
||||
|
||||
inter = max(0, inter_x2 - inter_x1) * max(0, inter_y2 - inter_y1)
|
||||
union = aw * ah + bw * bh - inter
|
||||
|
||||
return inter / union if union > 0 else 0
|
||||
|
||||
def iou_distance(self, tracks, detections):
|
||||
"""Compute IoU distance matrix"""
|
||||
iou_mat = np.zeros((len(tracks), len(detections)), dtype=np.float32)
|
||||
|
||||
for i, t in enumerate(tracks):
|
||||
for j, d in enumerate(detections):
|
||||
iou = self.iou(t.tlwh, d[:4])
|
||||
|
||||
# Penalize class mismatch
|
||||
if len(d) > 5 and t.class_id != int(d[5]):
|
||||
iou *= 0.3
|
||||
|
||||
# Boost for active tracks
|
||||
if t.time_since_update == 0:
|
||||
iou *= 1.1
|
||||
|
||||
iou_mat[i, j] = iou
|
||||
|
||||
return iou_mat
|
||||
|
||||
def linear_assignment(self, cost_matrix, thresh):
|
||||
"""Perform linear assignment with threshold"""
|
||||
if cost_matrix.size == 0:
|
||||
return np.empty((0, 2), dtype=int), tuple(range(cost_matrix.shape[0])), tuple(range(cost_matrix.shape[1]))
|
||||
|
||||
matches, unmatched_a, unmatched_b = [], [], []
|
||||
row_ind, col_ind = linear_sum_assignment(-cost_matrix)
|
||||
|
||||
for r, c in zip(row_ind, col_ind):
|
||||
if cost_matrix[r, c] < thresh:
|
||||
unmatched_a.append(r)
|
||||
unmatched_b.append(c)
|
||||
else:
|
||||
matches.append([r, c])
|
||||
|
||||
unmatched_a += list(set(range(cost_matrix.shape[0])) - set(row_ind))
|
||||
unmatched_b += list(set(range(cost_matrix.shape[1])) - set(col_ind))
|
||||
|
||||
return np.array(matches), tuple(unmatched_a), tuple(unmatched_b)
|
||||
|
||||
def update(self, detections):
|
||||
"""Two-stage association with stability fixes"""
|
||||
# Predict all tracks
|
||||
for t in self.tracks:
|
||||
t.predict()
|
||||
|
||||
confirmed_tracks = [t for t in self.tracks if t.hits >= 2]
|
||||
unconfirmed_tracks = [t for t in self.tracks if t.hits < 2]
|
||||
|
||||
# Separate high and low confidence detections
|
||||
high_dets = [d for d in detections if d[4] >= self.high_thresh]
|
||||
low_dets = [d for d in detections if self.low_thresh <= d[4] < self.high_thresh]
|
||||
|
||||
# First association
|
||||
unmatched_tracks = []
|
||||
unmatched_dets = []
|
||||
|
||||
if len(confirmed_tracks) > 0 and len(high_dets) > 0:
|
||||
iou_mat = self.iou_distance(confirmed_tracks, high_dets)
|
||||
matches, u_tracks, u_dets = self.linear_assignment(iou_mat, self.match_thresh)
|
||||
|
||||
for m in matches:
|
||||
track_idx, det_idx = m[0], m[1]
|
||||
class_id = int(high_dets[det_idx][5]) if len(high_dets[det_idx]) > 5 else 0
|
||||
class_name = self.get_class_name(class_id)
|
||||
confirmed_tracks[track_idx].update(high_dets[det_idx][:4],
|
||||
high_dets[det_idx][4],
|
||||
class_id,
|
||||
class_name)
|
||||
|
||||
unmatched_tracks = [confirmed_tracks[i] for i in u_tracks]
|
||||
unmatched_dets = [high_dets[i] for i in u_dets]
|
||||
else:
|
||||
unmatched_tracks = confirmed_tracks
|
||||
unmatched_dets = high_dets
|
||||
|
||||
# Second association with low confidence
|
||||
if len(unmatched_tracks) > 0 and len(low_dets) > 0:
|
||||
iou_mat = self.iou_distance(unmatched_tracks, low_dets)
|
||||
matches, u_tracks, _ = self.linear_assignment(iou_mat, self.match_thresh * 0.6)
|
||||
|
||||
for m in matches:
|
||||
track_idx, det_idx = m[0], m[1]
|
||||
class_id = int(low_dets[det_idx][5]) if len(low_dets[det_idx]) > 5 else 0
|
||||
class_name = self.get_class_name(class_id)
|
||||
unmatched_tracks[track_idx].update(low_dets[det_idx][:4],
|
||||
low_dets[det_idx][4],
|
||||
class_id,
|
||||
class_name)
|
||||
|
||||
unmatched_tracks = [unmatched_tracks[i] for i in u_tracks]
|
||||
|
||||
# Deal with unconfirmed tracks
|
||||
if len(unmatched_dets) > 0:
|
||||
for t in unconfirmed_tracks:
|
||||
iou_mat = self.iou_distance([t], unmatched_dets)
|
||||
if iou_mat.size > 0:
|
||||
best_match = np.argmax(iou_mat[0])
|
||||
if iou_mat[0, best_match] > self.match_thresh * 0.5:
|
||||
class_id = int(unmatched_dets[best_match][5]) if len(unmatched_dets[best_match]) > 5 else 0
|
||||
class_name = self.get_class_name(class_id)
|
||||
t.update(unmatched_dets[best_match][:4],
|
||||
unmatched_dets[best_match][4],
|
||||
class_id,
|
||||
class_name)
|
||||
unmatched_dets.pop(best_match)
|
||||
|
||||
# Create new tracks
|
||||
for d in unmatched_dets:
|
||||
if d[4] > self.track_thresh:
|
||||
overlaps = False
|
||||
for t in self.tracks:
|
||||
if self.iou(t.tlwh, d[:4]) > 0.3:
|
||||
overlaps = True
|
||||
break
|
||||
|
||||
if not overlaps:
|
||||
class_id = int(d[5]) if len(d) > 5 else 0
|
||||
class_name = self.get_class_name(class_id)
|
||||
track_id = self.get_next_id(class_id)
|
||||
self.tracks.append(Track(d[:4], d[4], track_id, class_id, class_name))
|
||||
|
||||
# Remove old tracks
|
||||
self.tracks = [t for t in self.tracks if t.time_since_update < self.buffer_size]
|
||||
|
||||
return self.tracks
|
||||
|
||||
|
||||
|
||||
# =========================================================
|
||||
# CLEAN CLASS NAME
|
||||
# =========================================================
|
||||
def clean_name(name):
|
||||
return name.replace("LEFT_", "").replace("RIGHT_", "").replace("_Start", "").replace("_Stop", "")
|
||||
|
||||
# =========================================================
|
||||
# PARSE JSON
|
||||
# =========================================================
|
||||
def parse_json(path):
|
||||
with open(path, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
try:
|
||||
data = json.loads(content) # proper JSON
|
||||
except json.JSONDecodeError:
|
||||
data = ast.literal_eval(content) # handles single quotes etc.
|
||||
|
||||
parsed = {}
|
||||
|
||||
for asset in data["Assets"]:
|
||||
cls = clean_name(asset[0])
|
||||
frame = int(asset[2])
|
||||
|
||||
x1, y1 = asset[3]
|
||||
x2, y2 = asset[4]
|
||||
|
||||
if frame not in parsed:
|
||||
parsed[frame] = []
|
||||
|
||||
parsed[frame].append({
|
||||
"class": cls,
|
||||
"bbox": [x1, y1, x2, y2],
|
||||
"id": asset[1]
|
||||
})
|
||||
|
||||
return parsed
|
||||
|
||||
# =========================================================
|
||||
# SAM
|
||||
# =========================================================
|
||||
def load_sam():
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
sam = sam_model_registry[MODEL_TYPE](checkpoint=SAM_CHECKPOINT)
|
||||
sam.to(device)
|
||||
return SamPredictor(sam)
|
||||
|
||||
def get_mask(image, bbox, predictor):
|
||||
masks, scores, _ = predictor.predict(
|
||||
box=np.array(bbox),
|
||||
multimask_output=True
|
||||
)
|
||||
return masks[np.argmax(scores)]
|
||||
|
||||
# =========================================================
|
||||
# TEMPORAL FRAMES"""
|
||||
Reads YOLO label--> convert to pixel format --> Feed the box to Sam ---> Get mask
|
||||
"""
|
||||
# =========================================================
|
||||
def get_frames(center, total):
|
||||
frames = []
|
||||
for i in range(-NUM_FRAMES, NUM_FRAMES + 1):
|
||||
f = center + i * FRAME_GAP
|
||||
if 0 <= f < total:
|
||||
frames.append(f)
|
||||
return frames
|
||||
|
||||
# =========================================================
|
||||
# MASK → POLYGON
|
||||
# =========================================================
|
||||
def mask_to_polygon(mask):
|
||||
mask_uint8 = (mask * 255).astype(np.uint8)
|
||||
contours, _ = cv2.findContours(mask_uint8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
|
||||
|
||||
polygons = []
|
||||
for cnt in contours:
|
||||
if len(cnt) >= 3:
|
||||
poly = cnt.squeeze().tolist()
|
||||
polygons.append(poly)
|
||||
|
||||
return polygons
|
||||
|
||||
# =========================================================
|
||||
# OVERLAY MASK
|
||||
# =========================================================
|
||||
def overlay_mask(image, mask):
|
||||
overlay = image.copy()
|
||||
color = np.array([255, 0, 0]) # RED
|
||||
|
||||
overlay[mask == 1] = overlay[mask == 1] * 0.5 + color * 0.5
|
||||
return overlay.astype(np.uint8)
|
||||
|
||||
# =========================================================
|
||||
# MAIN PIPELINE
|
||||
# =========================================================
|
||||
def process_video():
|
||||
predictor = load_sam()
|
||||
parsed = parse_json(JSON_PATH)
|
||||
|
||||
cap = cv2.VideoCapture(VIDEO_PATH)
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
|
||||
bbox_json = {}
|
||||
poly_json = {}
|
||||
|
||||
for center_frame in parsed:
|
||||
|
||||
temporal_frames = get_frames(center_frame, total_frames)
|
||||
|
||||
for obj_idx, obj in enumerate(parsed[center_frame]):
|
||||
cls = obj["class"]
|
||||
init_bbox = obj["bbox"]
|
||||
obj_id = obj["id"]
|
||||
|
||||
tracker = ByteTrackEnhanced(class_names=[cls])
|
||||
|
||||
x1, y1, x2, y2 = init_bbox
|
||||
tlwh = [x1, y1, x2-x1, y2-y1]
|
||||
|
||||
init_det = np.array([[*tlwh, 0.9, 0]])
|
||||
tracker.update(init_det)
|
||||
|
||||
for f in temporal_frames:
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, f)
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
continue
|
||||
|
||||
detections = init_det if f == center_frame else np.empty((0, 6))
|
||||
tracks = tracker.update(detections)
|
||||
|
||||
if len(tracks) == 0:
|
||||
continue
|
||||
|
||||
t = tracks[0]
|
||||
x, y, w, h = t.tlwh
|
||||
bbox = [int(x), int(y), int(x+w), int(y+h)]
|
||||
|
||||
image_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
predictor.set_image(image_rgb)
|
||||
|
||||
mask = get_mask(image_rgb, bbox, predictor)
|
||||
mask_uint8 = (mask * 255).astype(np.uint8)
|
||||
|
||||
# SAVE MASK
|
||||
cv2.imwrite(os.path.join(MASK_DIR, f"{f}.png"), mask_uint8)
|
||||
|
||||
# SAVE MASKED IMAGE
|
||||
extracted = cv2.bitwise_and(image_rgb, image_rgb, mask=mask_uint8)
|
||||
cv2.imwrite(os.path.join(EXTRACT_DIR, f"{f}.png"),
|
||||
cv2.cvtColor(extracted, cv2.COLOR_RGB2BGR))
|
||||
|
||||
# OVERLAY IMAGE
|
||||
overlay = overlay_mask(image_rgb, mask)
|
||||
cv2.imwrite(os.path.join(IMG_DIR, f"{f}.png"),
|
||||
cv2.cvtColor(overlay, cv2.COLOR_RGB2BGR))
|
||||
|
||||
# STORE BBOX JSON
|
||||
if str(f) not in bbox_json:
|
||||
bbox_json[str(f)] = {}
|
||||
|
||||
if cls not in bbox_json[str(f)]:
|
||||
bbox_json[str(f)][cls] = []
|
||||
|
||||
bbox_json[str(f)][cls].append([str(obj_id), [bbox[0], bbox[1]], [bbox[2], bbox[3]]])
|
||||
|
||||
# STORE POLYGON JSON
|
||||
polygons = mask_to_polygon(mask)
|
||||
|
||||
if str(f) not in poly_json:
|
||||
poly_json[str(f)] = {}
|
||||
|
||||
if cls not in poly_json[str(f)]:
|
||||
poly_json[str(f)][cls] = []
|
||||
|
||||
poly_json[str(f)][cls].append([str(obj_id), polygons])
|
||||
|
||||
print(f"[DONE] Frame {f}")
|
||||
|
||||
cap.release()
|
||||
|
||||
# SAVE JSON FILES
|
||||
with open(os.path.join(OUTPUT_DIR, "bboxes.json"), "w") as f:
|
||||
json.dump(bbox_json, f, indent=2)
|
||||
|
||||
with open(os.path.join(OUTPUT_DIR, "polygons.json"), "w") as f:
|
||||
json.dump(poly_json, f, indent=2)
|
||||
|
||||
|
||||
# =========================================================
|
||||
# ENTRY
|
||||
# =========================================================
|
||||
if __name__ == "__main__":
|
||||
process_video()
|
||||
Reference in New Issue
Block a user