CoolFace
Apppublic

Anmolkhurana88/Sign-Language-Translator

sourceHugging Faceunknownupdated 1mo agoView on Hugging Face
0likes
main.py161 linesDownload Raw Back to root
1from fastapi import FastAPI, WebSocket, WebSocketDisconnect, UploadFile, File
2from starlette.websockets import WebSocketState
3from fastapi.middleware.cors import CORSMiddleware
4import asyncio
5import time
6
7from landmark_extracter import extract_landmarks, process_frame, default_landmarks, correct_landmarks
8from word_level_model import load_word_model, predict_word_gloss
9from text_language_generator import generate_continue_text, GlossBuffer, create_text_buffer
10from frame_handler import FrameBuffer, decode_frame, decode_image_file
11import numpy as np
12
13app = FastAPI()
14
15# CORS Handler
16app.add_middleware(
17    CORSMiddleware,
18    allow_origins=["*"],
19    allow_credentials=True,
20    allow_methods=["*"],
21    allow_headers=["*"],
22)
23
24word_model = load_word_model('saved_models/word_level_model_states_include.pth')
25thres_word_conf = 0.8
26INFERENCE_INTERVAL = 0.2  # 0.2 seconds
27
28@app.get("/")
29async def root():
30    return {"status": "ok", "message": "Sign Language Translator is running on Hugging Face Spaces"}
31
32def gloss_prediction(frameData, frame_buffer, gloss_buffer, prev_lm, counter):
33    """Receives frames, predicts glosses, and fills buffer."""
34    try:
35        frame = decode_frame(frameData)
36
37        if frame is None:
38            return {"status": "error", "message": "Invalid frame data"}
39
40        # Extract and correct landmarks
41        curr_lm = extract_landmarks(frame)
42        corrected_lm = correct_landmarks(curr_lm, prev_lm)
43        prev_lm.update(curr_lm)
44
45        frame_buffer.add_frame(corrected_lm)
46
47        frame_seq = frame_buffer.get_frames()
48
49        if len(frame_seq) == 0 or time.time() - counter['last_inference_time'] < INFERENCE_INTERVAL:
50            # Not enough frames yet or recently predicted gloss
51            return None
52        
53        counter['last_inference_time'] = time.time()
54        
55        word_gloss, word_conf = predict_word_gloss(word_model, frame_seq)
56        print(f"Predicted Word Gloss: {word_gloss} with confidence {word_conf}")
57
58        if word_conf >= thres_word_conf:
59            gloss_buffer.append_gloss(word_gloss)
60
61        # text = generate_continue_text(text_buffer, gloss_buffer)
62        # text = word_gloss + " " + sentence_gloss
63
64        result = {
65            "word_confidence": word_conf,
66        }
67
68        return {"status": "success", "result": result}
69    
70    except Exception as e:
71        print(f"Error Processing Frame: {e}")
72        return {"status": "error", "message": str(e)}
73
74
75def text_generation(gloss_buffer, text_buffer, counter):
76    """Reads glosses and generates text asynchronously."""
77    try:
78        gen_text = generate_continue_text(gloss_buffer, text_buffer, counter)
79
80        if gen_text:
81            text_buffer.extend(gen_text.split())
82
83        return {"status": "success", "result": {"text": gen_text}}
84
85    except Exception as e:
86        print(f"Error Generating Text: {e}")
87        return {"status": "error", "message": str(e)}
88
89
90@app.websocket("/ws")
91async def websocket_endpoint(websocket: WebSocket):
92    await websocket.accept()
93    print("Client connected")
94
95    frame_buffer = FrameBuffer(max_size=20)
96    gloss_buffer = GlossBuffer()
97    text_buffer = create_text_buffer()
98
99    prev_lm = default_landmarks.copy()
100
101    counter = {'last_inference_time': time.time(), 'last_text_time': time.time()}
102
103    async def gloss_prediction_loop(websocket, frame_buffer, gloss_buffer):
104        async for frame_data in websocket.iter_text():
105            res = gloss_prediction(frame_data, frame_buffer, gloss_buffer, prev_lm, counter)
106
107            if res is not None:
108                await websocket.send_json(res)
109
110        # print(len(frame_buffer.buffer), "frames in buffer")
111        # pose_npy = np.array(frame_buffer.get_frames())
112        # pose_npy.dump("buffered_frames.npy")
113        # print("Saved buffered frames to buffered_frames.npy")
114
115    async def text_generation_loop(websocket, gloss_buffer, text_buffer):
116        while websocket.client_state == WebSocketState.CONNECTED:
117            await asyncio.sleep(1.5)  # Check every second
118
119            if websocket.client_state == WebSocketState.DISCONNECTED:
120                break
121            
122            res = text_generation(gloss_buffer, text_buffer, counter)
123
124            if websocket.client_state == WebSocketState.DISCONNECTED:
125                break
126
127            await websocket.send_json(res)
128
129    producer_task = asyncio.create_task(gloss_prediction_loop(websocket, frame_buffer, gloss_buffer))
130    consumer_task = asyncio.create_task(text_generation_loop(websocket, gloss_buffer, text_buffer))
131
132    try:
133        await asyncio.gather(producer_task, consumer_task)
134    except WebSocketDisconnect:
135        print("Client disconnected")
136    finally:
137        producer_task.cancel()
138        consumer_task.cancel()
139
140
141@app.post('/upload-image')
142async def upload_image(image: UploadFile = File(...)):
143    try:
144        # content = await image.read()
145        
146        # frame = decode_image_file(content)
147
148        # landmarks = extract_landmarks(frame)
149        # word_gloss, word_conf = predict_word_gloss(word_model, landmarks)
150
151        # response = {
152        #     "text": word_gloss.lower() if word_conf >= thres_word_conf else "",
153        #     "word_confidence": word_conf,
154        # }
155
156        response = {"message": "Image upload endpoint is depreciated."}
157
158        return {"status": "success", "result": response}
159    
160    except Exception as e:
161        return {"status": "error", "message": str(e)}