Spaces:
Running
Running
Update app.py
Browse files
app.py
CHANGED
|
@@ -47,21 +47,11 @@ def sequence_to_wav(sequence: np.ndarray, sample_rate=16000) -> str:
|
|
| 47 |
|
| 48 |
|
| 49 |
@app.get("/generate")
|
| 50 |
-
def generate(
|
| 51 |
-
"""
|
| 52 |
-
Generate a music sequence and return as WAV file.
|
| 53 |
-
"""
|
| 54 |
with torch.no_grad():
|
| 55 |
-
|
| 56 |
-
try:
|
| 57 |
-
input_vec = music_model.text_to_embedding(prompt).to(DEVICE)
|
| 58 |
-
except AttributeError:
|
| 59 |
-
# fallback to random vector if text embedding not implemented
|
| 60 |
-
input_vec = torch.randn(1, music_model.hidden_dim).to(DEVICE)
|
| 61 |
-
else:
|
| 62 |
-
input_vec = torch.randn(1, music_model.hidden_dim).to(DEVICE)
|
| 63 |
-
|
| 64 |
-
output = music_model.generate(input_vec, seq_len=seq_len).squeeze(0).cpu().numpy()
|
| 65 |
|
| 66 |
-
|
| 67 |
-
|
|
|
|
|
|
|
|
|
| 47 |
|
| 48 |
|
| 49 |
@app.get("/generate")
|
| 50 |
+
def generate(seq_len: int = 64):
|
|
|
|
|
|
|
|
|
|
| 51 |
with torch.no_grad():
|
| 52 |
+
tokens = music_model.generate(seq_len=seq_len, device=DEVICE)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 53 |
|
| 54 |
+
return {
|
| 55 |
+
"seq_len": seq_len,
|
| 56 |
+
"output": tokens.cpu().tolist()
|
| 57 |
+
}
|