Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 49 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -241,6 +241,55 @@ waveform = processor.decode(audio_codes)
torchaudio.save("tts.wav", waveform.cpu(), 24_000)
```

## Finetuning

To finetune on your own data, make use of the `ChatMessage` interface. This requires you to:

1. map your raw dataset rows into `list[ChatMessage]`
2. use the [`LFM2AudioChatMapper`](src/liquid_audio/data/mapper.py) to create a preprocessed dataset
3. train a model from the preprocessed dataset with `LFM2DataLoader`

First, install project dependencies:

```bash
uv sync
```

### Preprocess

Before training, convert dataset into our preprocessed training format.

To do that, define an iterator that yields one `list[ChatMessage]` per sample in your dataset.
The `LFM2AudioChatMapper` handles turning those messages into model-ready features.

`ChatMessage` supports:

* `TextSegment(text=...)`
* `AudioSegment(audio=...)`
* `InterleavedSegment(text=..., audio=...)`

See [examples/preprocess_jenny_tts.py](examples/preprocess_jenny_tts.py) for an example of how to preprocess the
[Jenny TTS Dataset](https://huggingface.co/datasets/reach-vb/jenny_tts_dataset) for TTS model finetuning.

Run preprocessing with:

```bash
python -m examples.preprocess_jenny_tts
```

This writes a preprocessed dataset to `data/jenny_tts/train`.

### Train

Training reads a preprocessed dataset.

For example, to finetune a model on the [Jenny TTS Dataset](https://huggingface.co/datasets/reach-vb/jenny_tts_dataset)
using the preprocessed dataset from before, run:

```bash
python -m examples.train
```


## License
The code in this repository and associated weights are licensed under the [LFM Open License v1.0](LICENSE).
Expand Down
45 changes: 45 additions & 0 deletions examples/preprocess_jenny_tts.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
from __future__ import annotations

from collections.abc import Iterator

from datasets import Audio, load_dataset

from liquid_audio import LFM2AudioProcessor
from liquid_audio.data.mapper import LFM2AudioChatMapper
from liquid_audio.data.preprocess import preprocess_dataset
from liquid_audio.data.types import AudioSegment, ChatMessage, TextSegment


class JennyTTSIterator:
def __init__(
self,
split: str = "train",
system_prompt: str = "Perform TTS. Use the Irish female voice.",
) -> None:
self.split = split
self.system_prompt = system_prompt

def __iter__(self) -> Iterator[list[ChatMessage]]:
ds = load_dataset("reach-vb/jenny_tts_dataset", split=self.split)
ds = ds.cast_column("audio", Audio(decode=False))

for row in ds:
text = row["transcription"]
audio = row["audio"]["bytes"]
yield [
ChatMessage(role="system", content=[TextSegment(text=self.system_prompt)]),
ChatMessage(role="user", content=[TextSegment(text=text)]),
ChatMessage(role="assistant", content=[AudioSegment(audio=audio)]),
]


if __name__ == "__main__":
processor = LFM2AudioProcessor.from_pretrained("LiquidAI/LFM2.5-Audio-1.5B", device="cuda").eval()
mapper = LFM2AudioChatMapper(processor)
data = JennyTTSIterator()
preprocess_dataset(
data=data,
output_path="data/jenny_tts/train",
mapper=mapper,
max_context_length=256, # skips 163 JennyTTS samples
)
45 changes: 45 additions & 0 deletions examples/train.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
from __future__ import annotations

import argparse
from pathlib import Path

from liquid_audio.data.dataloader import LFM2DataLoader
from liquid_audio.trainer import Trainer


def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--model-id", default="LiquidAI/LFM2.5-Audio-1.5B")
parser.add_argument("--data", default="data/jenny_tts/train")
parser.add_argument("--context-length", type=int, default=256)
parser.add_argument("--batch-size", type=int, default=64)
parser.add_argument("--max-steps", type=int, default=5000)
parser.add_argument("--warmup-steps", type=int, default=250)
parser.add_argument("--lr", type=float, default=1e-4)
parser.add_argument("--num-workers", type=int, default=8)
parser.add_argument("--output-dir", default="tmp")
return parser.parse_args()


if __name__ == "__main__":
args = parse_args()
dataset_path = Path(args.data)
if not dataset_path.exists():
raise FileNotFoundError(f"Preprocessed dataset not found at {dataset_path}. Run preprocessing before training.")

train_data = LFM2DataLoader(dataset_path=str(dataset_path), context_length=args.context_length)

trainer = Trainer(
model_id=args.model_id,
train_data=train_data,
lr=args.lr,
batch_size=args.batch_size,
max_steps=args.max_steps,
warmup_steps=args.warmup_steps,
dataloader_num_workers=args.num_workers,
logging_interval=10,
save_interval=500,
val_interval=100,
output_dir=args.output_dir,
)
trainer.train()
4 changes: 3 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[project]
name = "liquid-audio"
version = "1.1.0"
version = "1.2.0"
description = "Liquid Audio - Speech-to-Speech audio models"
readme = "README.md"
authors = [
Expand All @@ -11,6 +11,7 @@ license-files = ["LICENSE"]
requires-python = ">=3.12"
dependencies = [
"accelerate>=1.10.1",
"datasets>=4.8.4",
"einops>=0.8.1",
"librosa>=0.11.0",
"sentencepiece>=0.2.1",
Expand Down Expand Up @@ -125,6 +126,7 @@ check_untyped_defs = true
module = [
"accelerate.*",
"torchaudio.*",
"datasets.*",
]
ignore_missing_imports = true

Expand Down
Empty file.
75 changes: 75 additions & 0 deletions src/liquid_audio/data/dataloader.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
from __future__ import annotations

from pathlib import Path

import torch
import torch.nn.functional as F
from datasets import Dataset, load_from_disk
from torch.utils.data import Dataset as TorchDataset

from liquid_audio.data.types import LFM2AudioModelInput, LFM2AudioRow
from liquid_audio.utils import LFMModality


class LFM2DataLoader(TorchDataset[LFM2AudioRow]):
def __init__(
self,
dataset_path: str,
context_length: int = 4096,
) -> None:
self.dataset_path = Path(dataset_path)
self.context_length = context_length
self.dataset: Dataset = load_from_disk(self.dataset_path)

def __len__(self) -> int:
return len(self.dataset)

def __getitem__(self, idx: int) -> LFM2AudioRow:
row = self.dataset[idx]

text = torch.as_tensor(row["text"], dtype=torch.long)
audio_in = torch.as_tensor(row["audio_in"], dtype=torch.float32)
audio_in_lens = torch.as_tensor(row["audio_in_lens"], dtype=torch.long)
audio_out = torch.as_tensor(row["audio_out"], dtype=torch.long)
modality = torch.as_tensor(row["modality_flag"], dtype=torch.long)
supervision = torch.as_tensor(row["supervision_mask"], dtype=torch.bool)

pad_len = self.context_length - int(modality.shape[1])
if pad_len < 0:
raise ValueError(
f"sample at index {idx} has {modality.shape[1]} tokens, "
f"which is longer than context_length={self.context_length}"
)

text = F.pad(text, (0, pad_len))
modality = F.pad(modality, (0, pad_len), value=int(LFMModality.TEXT))
supervision = F.pad(supervision, (0, pad_len), value=False)

return LFM2AudioRow(
text=text,
audio_in=audio_in,
audio_in_lens=audio_in_lens,
audio_out=audio_out,
modality_flag=modality,
supervision_mask=supervision,
)


def lfm2_collator(batch: list[LFM2AudioRow]) -> LFM2AudioModelInput:
audio_in = torch.cat([row.audio_in for row in batch], dim=1)
audio_in_lens = torch.cat([row.audio_in_lens for row in batch], dim=0)

text = torch.cat([row.text for row in batch], dim=1)
audio_out = torch.cat([row.audio_out for row in batch], dim=1)

modality_flag = torch.cat([row.modality_flag for row in batch], dim=0)
supervision_mask = torch.cat([row.supervision_mask for row in batch], dim=0)

return LFM2AudioModelInput(
text=text,
audio_in=audio_in,
audio_in_lens=audio_in_lens,
audio_out=audio_out,
modality_flag=modality_flag,
supervision_mask=supervision_mask,
)
Loading
Loading