Spaces:
Sleeping
Sleeping
| 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") | |