fomext commited on
Commit
70fccfa
·
verified ·
1 Parent(s): 8b80b8c

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +24 -8
app.py CHANGED
@@ -14,14 +14,14 @@ print("Audiocraft version:", audiocraft.__version__)
14
 
15
  app = FastAPI()
16
 
17
- MODEL_NAME = "facebook/musicgen-medium"
18
  OUTPUT_DIR = "outputs"
19
  os.makedirs(OUTPUT_DIR, exist_ok=True)
20
 
21
  print("Loading MusicGen model...")
22
  model = MusicGen.get_pretrained(MODEL_NAME)
23
  model.set_generation_params(
24
- duration=30,
25
  temperature=1.0,
26
  top_k=250,
27
  top_p=0.0
@@ -42,21 +42,29 @@ class JobStatus:
42
  FAILED = "failed"
43
 
44
 
 
 
 
 
 
 
 
 
45
  def worker():
46
- """Background worker that processes queued jobs"""
47
  while True:
48
  job_id = job_queue.get()
49
  job = jobs.get(job_id)
50
 
51
- if not job:
52
- job_queue.task_done()
53
- continue
54
-
55
  try:
56
  jobs[job_id]["status"] = JobStatus.PROCESSING
57
 
 
 
 
58
  wav = model.generate([job["prompt"]])[0]
59
 
 
 
60
  filename = f"{job_id}.wav"
61
  path = os.path.join(OUTPUT_DIR, filename)
62
 
@@ -73,13 +81,21 @@ def worker():
73
  "file_path": path
74
  })
75
 
 
 
 
 
 
 
76
  except Exception as e:
77
  jobs[job_id].update({
78
  "status": JobStatus.FAILED,
79
  "error": str(e)
80
  })
81
 
82
- job_queue.task_done()
 
 
83
 
84
 
85
  # Start worker thread
 
14
 
15
  app = FastAPI()
16
 
17
+ MODEL_NAME = "facebook/musicgen-small"
18
  OUTPUT_DIR = "outputs"
19
  os.makedirs(OUTPUT_DIR, exist_ok=True)
20
 
21
  print("Loading MusicGen model...")
22
  model = MusicGen.get_pretrained(MODEL_NAME)
23
  model.set_generation_params(
24
+ duration=10,
25
  temperature=1.0,
26
  top_k=250,
27
  top_p=0.0
 
42
  FAILED = "failed"
43
 
44
 
45
+ import signal
46
+
47
+ class GenerationTimeout(Exception):
48
+ pass
49
+
50
+ def timeout_handler(signum, frame):
51
+ raise GenerationTimeout()
52
+
53
  def worker():
 
54
  while True:
55
  job_id = job_queue.get()
56
  job = jobs.get(job_id)
57
 
 
 
 
 
58
  try:
59
  jobs[job_id]["status"] = JobStatus.PROCESSING
60
 
61
+ signal.signal(signal.SIGALRM, timeout_handler)
62
+ signal.alarm(420) # 7 minutes max
63
+
64
  wav = model.generate([job["prompt"]])[0]
65
 
66
+ signal.alarm(0)
67
+
68
  filename = f"{job_id}.wav"
69
  path = os.path.join(OUTPUT_DIR, filename)
70
 
 
81
  "file_path": path
82
  })
83
 
84
+ except GenerationTimeout:
85
+ jobs[job_id].update({
86
+ "status": JobStatus.FAILED,
87
+ "error": "Generation timed out (CPU limit)"
88
+ })
89
+
90
  except Exception as e:
91
  jobs[job_id].update({
92
  "status": JobStatus.FAILED,
93
  "error": str(e)
94
  })
95
 
96
+ finally:
97
+ job_queue.task_done()
98
+
99
 
100
 
101
  # Start worker thread