Stable Audio
Stable Audio
Stable Audio는 Stability AI가 공개한 텍스트-오디오 생성(Text-to-Audio) 모델과 그 학습·추론 코드(training and inference code)를 담은 오픈소스 프로젝트예요. 자연어 프롬프트를 입력하면 그에 맞는 오디오를 생성하고, 학습용 코드를 통해 자신만의 오디오 생성 모델을 학습하거나 파인튜닝할 수도 있어요.
핵심 기술은 잠재 확산(latent diffusion) 방식이에요. 오디오를 잠재 공간(latent space)으로 압축한 뒤 확산 모델(diffusion model)로 점진적으로 노이즈를 제거하면서 음성을 복원하는 방식이죠. 여기서 쓰이는 오디오 압축기(오토인코더 프리트랜스폼)가 바로 DAC(Descript Audio Codec) 계열로, 모델은 오디오를 고품질로 잠재 표현으로 변환하고 다시 원래 신호로 복원해요.
출처: 문서
본문
Stable Audio는 텍스트 프롬프트를 바탕으로 오디오를 생성하는 모델이에요. AudioLDM 등과 같은 latent diffusion 아키텍처를 기반으로 하며, 사전 학습된 autoencoder(pretransform)가 오디오를 잠재 표현으로 압축하면 확산 모델이 조건(텍스트)을 반영해 그 잠재 표현에서 오디오를 생성해요. 이 압축 단계에 사용되는 게 DAC 계열 코덱 기반의 오토인코더이고요.
프로젝트는 설정 파일 기반으로 동작해요. 모델의 하이퍼파라미터와 학습·추론 정보를 담은 모델 설정(model config), 그리고 학습 데이터셋 정보를 담은 데이터셋 설정(dataset config) 이 JSON 파일로 정의되고요. 모델 타입은 autoencoder, diffusion_uncond, diffusion_cond, diffusion_cond_inpaint, diffusion_autoencoder, lm 등이 지원돼요.
설치
PyTorch 2.5 이상이 필요하고, 개발은 Python 3.10에서 이루어져요. 의존성 관리는 빠르고 재현 가능한 uv를 사용해요. pip install uv로 설치하거나 uv docs를 참고하면 돼요.
저장소를 클론하고 의존성을 설치해요:
$ git clone https://github.com/Stability-AI/stable-audio-tools.git
$ cd stable-audio-tools
# Inference only
$ uv sync
# Training
$ uv sync --extra train
# Everything (training + Gradio UI)
$ uv sync --extra train --extra ui
성능을 위해 Flash Attention 설치를 권장해요. uv sync 이후 Flash Attention 저장소의 설치 안내를 따라 설치하면 되고요.
스크립트는 uv run으로 실행해요:
$ uv run python train.py --dataset-config /path/to/config ...
$ uv run python run_gradio.py --pretrained-name stabilityai/stable-audio-open-1.0
또는 pip을 직접 써도 돼요:
$ pip install "stable-audio-tools[train]"
사용 예시 (Gradio 인터페이스)
훈련된 모델을 테스트할 수 있는 기본 Gradio 인터페이스를 제공해요. Hugging Face에서 stable-audio-open-1.0 모델의 이용 약관에 동의했다면 다음 명령으로 인터페이스를 띄울 수 있어요:
$ python3 ./run_gradio.py --pretrained-name stabilityai/stable-audio-open-1.0
run_gradio.py 는 다음과 같은 명령줄 인자를 받아요:
--pretrained-name: Hugging Face 저장소 이름. 모델 설정/체크포인트 대신 사전 학습 모델을 쓸 때 사용해요.--model-config: 로컬 모델의 설정 파일 경로--ckpt-path: 로컬 모델의 unwrapped 체크포인트 경로--pretransform-ckpt-path: 파인튜닝된 디코더 테스트용 pretransform 체크포인트 (선택)--share: true면 Gradio 데모의 공개 공유 링크 생성 (선택)--username,--password: Gradio 데모 로그인 설정 (선택)--model-half: true면 모델 가중치를 half-precision으로 (선택)
학습 (Training)
학습을 시작하려면 모델 설정 파일과 데이터셋 설정 파일이 필요해요. 학습 로그와 데모를 기록하려면 Weights & Biases 계정도 필요하고요:
$ wandb login
학습은 저장소 루트에서 train.py 를 실행해 시작해요:
$ python3 ./train.py --dataset-config /path/to/dataset/config --model-config /path/to/model/config --name harmonai_train
--name 은 Weights and Biases 런의 프로젝트 이름을 정해요.
학습 중 모델은 discrminator, EMA 복사본, 옵티마이저 상태 등을 담은 "training wrapper"(pl.LightningModule)로 감싸여요. 체크포인트에는 이 wrapper가 포함돼 크기가 커지는데, unwrap_model.py 를 실행하면 모델만 담은 체크포인트를 새로 저장할 수 있어요:
$ python3 ./unwrap_model.py --model-config /path/to/model/config --ckpt-path /path/to/wrapped/ckpt --name model_unwrap
Unwrapped 체크포인트는 추론 스크립트, 다른 모델의 pretransform으로 사용(예: latent diffusion용 오토인코더), 파인튜닝 등에 필요해요.
파인튜닝은 사전 학습 체크포인트에서 학습을 이어가는 방식이에요. wrapped 체크포인트는 --ckpt-path 로, 사전 학습된 unwrapped 모델로 새로 시작하려면 --pretrained-ckpt-path 로 전달하면 돼요.
train.py 의 주요 추가 플래그는 다음과 같아요:
--config-file: 저장소 루트의 defaults.ini 경로 (다른 디렉터리에서 실행할 때 필요)--pretransform-ckpt-path: latent diffusion 등에서 사전 학습된 autoencoder 로드용 (unwrapped 체크포인트 필요)--save-dir: 체크포인트 저장 디렉터리--checkpoint-every: 체크포인트 저장 사이의 스텝 수 (기본값: 10000)--batch-size: GPU당 샘플 수. VRAM이 허용하는 만큼 크게 (기본값: 8)--num-gpus: 노드당 GPU 수 (기본값: 1)--num-nodes: 사용할 GPU 노드 수 (기본값: 1)--accum-batches: 작은 GPU에서 배치 크기를 키울 때 쓰는 gradient accumulation 배치 수--strategy: 분산 학습용 멀티 GPU 전략.deepspeed로 설정하면 DeepSpeed ZeRO Stage 2 사용 (기본값:--num_gpus> 1일 때ddp, 아니면 None)--precision: 학습에 사용할 부동소수점 정밀도 (기본값: 16)--num-workers: 데이터 로더가 쓰는 CPU 워커 수--seed: PyTorch의 RNG 시드. 결정적 학습에 도움
설정 (Configurations)
학습·추론 코드는 모델 하이퍼파라미터, 학습 설정, 데이터셋 정보를 정의하는 JSON 설정 파일 기반으로 동작해요.
- 모델 설정(model config): 모델을 로드하는 데 필요한 모든 정보를 담아요. 최상위에
model_type(모델 타입),sample_size(학습에 제공하는 오디오 길이, 샘플 단위),sample_rate(오디오 샘플 레이트, Hz),audio_channels(오디오 채널 수, 기본 2 / 모노는 1),model(모델 타입별 설정),training(학습·데모 설정) 등을 정의해요. - 데이터셋 설정(dataset config): 로컬 오디오 파일 디렉터리와 Amazon S3에 저장된 WebDataset 두 종류의 데이터 소스를 지원해요. 자세한 내용은 데이터셋 설정 문서에서 볼 수 있어요.