|
26 | 26 | sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'common')) |
27 | 27 | from grpc_auth import get_auth_interceptors |
28 | 28 | from model_utils import resolve_model_reference |
| 29 | +from audio_utils import write_pcm_wav |
29 | 30 |
|
30 | 31 |
|
31 | 32 | # Import dynamic loader for pipeline discovery |
@@ -883,6 +884,49 @@ def GenerateImage(self, request, context): |
883 | 884 |
|
884 | 885 | return backend_pb2.Result(message="Media generated", success=True) |
885 | 886 |
|
| 887 | + def SoundGeneration(self, request, context): |
| 888 | + if not request.dst: |
| 889 | + return backend_pb2.Result(success=False, message="request.dst is required") |
| 890 | + |
| 891 | + prompt = request.text or request.caption |
| 892 | + if not prompt: |
| 893 | + return backend_pb2.Result(success=False, message="request.text is required") |
| 894 | + |
| 895 | + try: |
| 896 | + generation_options = dict(self.options) |
| 897 | + if "num_inference_steps" in generation_options: |
| 898 | + generation_options["num_inference_steps"] = int( |
| 899 | + generation_options["num_inference_steps"] |
| 900 | + ) |
| 901 | + generation_options["prompt"] = prompt |
| 902 | + if request.HasField("duration"): |
| 903 | + generation_options["audio_length_in_s"] = request.duration |
| 904 | + if request.HasField("temperature"): |
| 905 | + generation_options["guidance_scale"] = request.temperature |
| 906 | + |
| 907 | + generated = self.pipe(**generation_options) |
| 908 | + if not hasattr(generated, "audios") or len(generated.audios) == 0: |
| 909 | + return backend_pb2.Result( |
| 910 | + success=False, |
| 911 | + message="The diffusers pipeline returned no audio", |
| 912 | + ) |
| 913 | + |
| 914 | + samples = generated.audios[0] |
| 915 | + if hasattr(samples, "reshape"): |
| 916 | + samples = samples.reshape(-1) |
| 917 | + if hasattr(samples, "tolist"): |
| 918 | + samples = samples.tolist() |
| 919 | + |
| 920 | + sampling_rate = getattr( |
| 921 | + getattr(getattr(self.pipe, "vae", None), "config", None), |
| 922 | + "sampling_rate", |
| 923 | + 16000, |
| 924 | + ) |
| 925 | + write_pcm_wav(request.dst, samples, sampling_rate) |
| 926 | + return backend_pb2.Result(success=True, message="Sound generated successfully") |
| 927 | + except Exception as err: |
| 928 | + return backend_pb2.Result(success=False, message=f"SoundGeneration error: {err}") |
| 929 | + |
886 | 930 | def UpscaleImage(self, request, context): |
887 | 931 | try: |
888 | 932 | if not request.src: |
|
0 commit comments