eBIM-Satu2 / app.py
Azerofth's picture
Update app.py
b2112e2 verified
Raw
History Blame Contribute Delete
5.6 kB
import gradio as gr
import matplotlib
matplotlib.use('Agg')
import cv2
import os
import tempfile
import base64
import numpy as np
import torch
import mediapipe as mp
from model import CustomLSTM, STGCNModel, CTRGCNModel, SkateFormerModel
from utils import mediapipe_detection, draw_styled_landmarks, extract_keypoints, get_adjacency_matrix, reshape_for_stgcn
# --- Setup ---
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
gloss = np.load('gloss.npy')
# Load your preferred model (Defaulting to SkateFormer based on your inference.py)
model = SkateFormerModel(num_classes=len(gloss)).to(device)
model.load_state_dict(torch.load('best_model.pth', map_location=device))
model.eval()
mp_holistic = mp.solutions.holistic
def process_video(video_path):
if video_path is None:
return None, "⚠️ Please record a video for prediction."
try:
cap = cv2.VideoCapture(video_path)
fps = cap.get(cv2.CAP_PROP_FPS) or 30
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
frames_for_display = []
sequence = []
with mp_holistic.Holistic(min_detection_confidence=0.5, min_tracking_confidence=0.5) as holistic:
while cap.isOpened():
ret, frame = cap.read()
if not ret:
break
# Extract keypoints using your existing logic
image, results = mediapipe_detection(frame, holistic)
draw_styled_landmarks(image, results)
keypoints = extract_keypoints(results)
sequence.append(keypoints)
frames_for_display.append(image)
cap.release()
# Padding/Sampling logic to ensure exactly 30 frames
if len(sequence) < 30:
sequence.extend([np.zeros(258)] * (30 - len(sequence)))
else:
# Uniformly sample 30 frames
idx = np.linspace(0, len(sequence) - 1, 30).astype(int)
sequence = [sequence[i] for i in idx]
# Model Inference
input_data = np.expand_dims(sequence, axis=0)
input_processed = reshape_for_stgcn(input_data)
input_tensor = torch.tensor(input_processed, dtype=torch.float32).to(device)
with torch.no_grad():
res = model(input_tensor)
prob = torch.nn.functional.softmax(res, dim=1)
confidence, max_idx = torch.max(prob, dim=1)
label = gloss[max_idx.item()]
result_text = f"{label} ({confidence.item() * 100:.2f}%)"
# Create a temporary file to save the video
fd, output_path = tempfile.mkstemp(suffix='.mp4')
os.close(fd)
# Use MP4V for mp4 files; ensures browser compatibility
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
out = cv2.VideoWriter(output_path, fourcc, fps, (width, height))
for frame in frames_for_display:
# Draw background rectangle for readability
cv2.rectangle(frame, (0, 0), (width, 60), (245, 117, 16), -1)
# Put prediction text on frame
cv2.putText(frame, result_text, (10, 45),
cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 255, 255), 2, cv2.LINE_AA)
out.write(frame)
out.release()
return output_path, "✅ Analysis Complete"
except Exception as e:
return None, f"❌ Error: {str(e)}"
def get_img_html(img_path):
"""
Reads a local image file and converts it to a Base64 HTML string.
This prevents "broken image" errors in Gradio.
"""
try:
with open(img_path, "rb") as image_file:
encoded_string = base64.b64encode(image_file.read()).decode()
# Return the HTML img tag with the data embedded
return f'<img src="data:image/png;base64,{encoded_string}" width="40" style="display: inline-block; margin-right: 10px; vertical-align: bottom;" />'
except Exception as e:
print(f"Could not load icon: {e}")
return "" # Return empty string if file missing
def reset_state():
return None, None, "Ready"
# --- Gradio UI ---
with gr.Blocks(title="eBIM-Satu (Beta Version)") as demo:
icon_html = get_img_html("favicon.png")
gr.Markdown(f"# {icon_html} eBIM-Satu")
gr.Markdown(f"### Malaysia Isolated Sign Language Recognition")
gr.Markdown("Click 'Record' to start. The system will process the sign after 3 seconds.")
with gr.Row():
with gr.Column(scale=1):
system_status = gr.Textbox(label="System Status", value="Ready", lines=1)
gr.Markdown(
"### Instructions\n1. Open your camera.\n2. Click the record button in the video box.\n3. Perform the sign clearly.\n4. Click 'Start Prediction'.")
with gr.Column(scale=2):
input_video = gr.Video(sources=["webcam"], format="mp4", label="Sign Camera", webcam_options=gr.WebcamOptions(mirror=False))
with gr.Row():
predict_btn = gr.Button("Start Prediction / Analyze", variant="primary", scale=2)
clear_btn = gr.Button("Clear / Reset", variant="secondary", scale=1)
with gr.Column(scale=2):
output_video = gr.Video(autoplay=True, show_label=False)
predict_btn.click(
fn=process_video,
inputs=input_video,
outputs=[output_video, system_status]
)
clear_btn.click(
fn=reset_state,
inputs=None,
outputs=[input_video, output_video, system_status]
)
demo.launch(theme=gr.themes.Soft(), favicon_path="favicon.png")