Anmolkhurana88/Sign-Language-Translator
0
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)}