uriva commited on
Commit
6d31c91
·
1 Parent(s): f7299a1

Add API-friendly endpoint bypassing FileData validation

Browse files
Files changed (1) hide show
  1. app.py +85 -0
app.py CHANGED
@@ -200,6 +200,64 @@ else:
200
  process_video_ui = _process_video_impl
201
 
202
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
203
  def extract_first_frame_thumbnail(
204
  video_path, output_path, size=(200, 200), force=False
205
  ):
@@ -384,6 +442,33 @@ with gr.Blocks(title="AutoGaze Demo", delete_cache=(86400, 86400)) as demo:
384
  ],
385
  ).then(fn=cleanup_gpu, inputs=None, outputs=None)
386
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
387
  # Clean up GPU memory when user disconnects
388
  demo.unload(cleanup_gpu)
389
 
 
200
  process_video_ui = _process_video_impl
201
 
202
 
203
+ def _process_video_api_impl(
204
+ file_path, gazing_ratio, task_loss_requirement, output_fps, progress=None
205
+ ):
206
+ """API-friendly endpoint that takes a file path string instead of gr.File.
207
+ Returns only the gazing video path for programmatic use."""
208
+ if not file_path or not os.path.exists(file_path):
209
+ return "Error: file not found"
210
+
211
+ metadata = extract_metadata(file_path)
212
+ if metadata[1] is None:
213
+ return "Error: could not read file"
214
+
215
+ _, tmp_path, total_frames, fps, _, _ = metadata
216
+
217
+ yield "Loading model..."
218
+
219
+ if progress:
220
+ progress(0.0, desc="Loading model...")
221
+ setup = get_model(device)
222
+
223
+ yield "Processing video..."
224
+
225
+ if progress:
226
+ progress(0.1, desc="Processing video...")
227
+
228
+ def update_progress(pct, msg):
229
+ if progress:
230
+ progress(pct, desc=msg)
231
+
232
+ model_gazing_ratio = gazing_ratio * (196 / 265)
233
+
234
+ for results in process_video(
235
+ tmp_path,
236
+ setup,
237
+ gazing_ratio=model_gazing_ratio,
238
+ task_loss_requirement=task_loss_requirement,
239
+ progress_callback=update_progress,
240
+ spatial_batch_size=2,
241
+ ):
242
+ yield "Processing..."
243
+
244
+ yield "Saving output..."
245
+
246
+ fps_to_use = output_fps if output_fps is not None else results["fps"]
247
+
248
+ gazing_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
249
+ save_video(results["gazing_frames"], gazing_file.name, fps_to_use)
250
+ gazing_file.close()
251
+
252
+ yield gazing_file.name
253
+
254
+
255
+ if ZEROGPU_AVAILABLE:
256
+ process_video_api = spaces.GPU(duration=120)(_process_video_api_impl)
257
+ else:
258
+ process_video_api = _process_video_api_impl
259
+
260
+
261
  def extract_first_frame_thumbnail(
262
  video_path, output_path, size=(200, 200), force=False
263
  ):
 
442
  ],
443
  ).then(fn=cleanup_gpu, inputs=None, outputs=None)
444
 
445
+ # --- API-friendly endpoint (hidden tab, bypasses FileData validation) ---
446
+ with gr.Tab("API", visible=False):
447
+ api_file_path = gr.Textbox(label="File Path")
448
+ api_gazing_ratio = gr.Slider(
449
+ minimum=round(1 / 196, 2),
450
+ maximum=round(265 / 196, 2),
451
+ step=0.01,
452
+ value=0.75,
453
+ label="Gazing Ratio",
454
+ )
455
+ api_task_loss = gr.Slider(
456
+ minimum=0.0,
457
+ maximum=1.5,
458
+ step=0.05,
459
+ value=0.7,
460
+ label="Task Loss Requirement",
461
+ )
462
+ api_output_fps = gr.Number(label="Output FPS", value=None)
463
+ api_button = gr.Button("Process (API)")
464
+ api_result = gr.Textbox(label="Result Path")
465
+ api_button.click(
466
+ fn=process_video_api,
467
+ inputs=[api_file_path, api_gazing_ratio, api_task_loss, api_output_fps],
468
+ outputs=[api_result],
469
+ api_name="process_video_api",
470
+ )
471
+
472
  # Clean up GPU memory when user disconnects
473
  demo.unload(cleanup_gpu)
474