Skip to content

CLI robustness: output path checked only after processing, errors go to stdout, GPU hard-coded to device 0 #21

Description

@joeljose

Severity: Medium. On the MIT samples, a CPU run takes over an hour (see the performance issue), so a failure at the very end is expensive.

Problems

  1. The output path is checked last. save_wav is the final step, so a directory that doesn't exist or can't be written raises FileNotFoundError/PermissionError as a traceback after all frames are processed, and the recovered audio is lost.
  2. Errors go to stdout. Every Error: / Warning: message uses plain print(). Scripts and pipelines can't separate diagnostics from output. (The VRAM warning is the only one sent to stderr.)
  3. GPU device 0 is hard-coded. torch.device('cuda'), get_device_name(0) and mem_get_info(0) are all fixed, and there's no --device flag, although the sibling projects have one.
  4. estimate_vram ignores its nlevels argument. It's a fixed 15× multiplier, and it only warns, so an OOM still happens later.
  5. No torch.no_grad(). The GPU forward pass runs without torch.inference_mode(). It's harmless today because the filters are buffers, but a future pytorch_wavelets could make them parameters and track gradients.
  6. Timing. Pipeline timing starts before validation, and the CPU path imports dtcwt twice.

Suggested fix

  • Resolve and check the output path at the start of main(): create the parent directory or fail, and do a test write with os.access.
  • Wrap the final write so a failure saves to a fallback path such as ./sound_<timestamp>.wav rather than losing the result.
  • Send all diagnostics to sys.stderr, or use logging. Add -q/--quiet.
  • Add --device cpu|cuda[:N], used for model placement, get_device_name and mem_get_info. It also allows CPU-torch runs in CI.
  • Use with torch.inference_mode(): around the forward and phase maths.
  • Either use nlevels in the estimate or drop the parameter. Catching OOM once and halving the batch automatically would be friendlier than exiting.

Acceptance criteria

  • An unwritable output path fails in under a second, before any frames are read.
  • Errors go to stderr.
  • --device is honoured.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    bugSomething isn't working

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions