mirror of
https://github.com/Nighthawk42/soprano-factory.git
synced 2026-08-30 04:30:21 +00:00
Refactored training pipeline to use YAML config and uv
- Move hyperparameters from hardcoded script values to `config.yaml` - Replace pip requirements with `pyproject.toml` and `uv` support (CUDA 12.6) - Refactor all scripts to use `pathlib` for robust path handling - Optimize `generate_dataset.py` with GPU acceleration - Register quantizer constants as buffers for proper device mapping - Update README with new installation and usage instructions
This commit is contained in:
@@ -4,3 +4,4 @@ test.py
|
||||
*.pth
|
||||
dist/
|
||||
*.egg-info/
|
||||
.vs
|
||||
@@ -33,6 +33,20 @@
|
||||
|
||||
## Installation
|
||||
|
||||
This project uses **[uv](https://github.com/astral-sh/uv)** for high-performance dependency management.
|
||||
|
||||
```bash
|
||||
git clone [https://github.com/ekwek1/soprano-factory.git](https://github.com/ekwek1/soprano-factory.git)
|
||||
cd soprano-factory
|
||||
|
||||
# Install dependencies (CPU)
|
||||
uv sync
|
||||
|
||||
# Install dependencies (CUDA 12.6)
|
||||
uv sync --extra gpu
|
||||
|
||||
## Manual Installation
|
||||
|
||||
```bash
|
||||
git clone https://github.com/ekwek1/soprano-factory.git
|
||||
cd soprano-factory
|
||||
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
# Hardware & Reproducibility
|
||||
device: "cuda:0"
|
||||
seed: 1337
|
||||
|
||||
# Learning Rate Schedule
|
||||
max_lr: 5.0e-4 # Use 5.0e-4 (with dot) to ensure it loads as a float
|
||||
warmup_ratio: 0.1
|
||||
cooldown_ratio: 0.1
|
||||
|
||||
# Training Dynamics
|
||||
batch_size: 4
|
||||
grad_accum_steps: 1
|
||||
seq_len: 1024
|
||||
max_steps: 10000
|
||||
val_freq: 250
|
||||
|
||||
# Optimizer & Model Config
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.1
|
||||
text_factor: 0.0 # Increase to train on text inputs
|
||||
+9
-2
@@ -1,10 +1,13 @@
|
||||
import json
|
||||
import pathlib
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
class AudioDataset(Dataset):
|
||||
def __init__(self, path):
|
||||
with open(path, encoding='utf-8') as f:
|
||||
# Convert string path to Path object if necessary for consistency
|
||||
self.path = pathlib.Path(path)
|
||||
with open(self.path, encoding='utf-8') as f:
|
||||
self.dataset = json.load(f)
|
||||
|
||||
def __len__(self):
|
||||
@@ -12,6 +15,10 @@ class AudioDataset(Dataset):
|
||||
|
||||
def __getitem__(self, idx):
|
||||
text, audio = self.dataset[idx]
|
||||
|
||||
# Format: [STOP][TEXT]<text prompt>[START]<audio tokens>[STOP]
|
||||
res = f"[STOP][TEXT]{text}[START]{''.join(list(map(lambda x: f'[{x}]', audio)))}[STOP]"
|
||||
# Optimization: Use a generator expression for joining tokens
|
||||
audio_tokens = ''.join(f'[{x}]' for x in audio)
|
||||
res = f"[STOP][TEXT]{text}[START]{audio_tokens}[STOP]"
|
||||
|
||||
return res
|
||||
+8
-36
@@ -14,10 +14,7 @@ def safe_log(x: torch.Tensor, clip_val: float = 5e-3) -> torch.Tensor:
|
||||
|
||||
|
||||
class SimpleMLP(nn.Module):
|
||||
def __init__(self,
|
||||
dim,
|
||||
intermediate_dim,
|
||||
):
|
||||
def __init__(self, dim, intermediate_dim):
|
||||
super().__init__()
|
||||
self.pwconv1 = nn.Linear(dim, intermediate_dim)
|
||||
self.act = nn.GELU()
|
||||
@@ -31,14 +28,7 @@ class SimpleMLP(nn.Module):
|
||||
|
||||
|
||||
class ConvNeXtBlock(nn.Module):
|
||||
"""ConvNeXt Block adapted from https://github.com/facebookresearch/ConvNeXt to 1D audio signal.
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
intermediate_dim (int): Dimensionality of the intermediate layer.
|
||||
layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling.
|
||||
Defaults to None.
|
||||
"""
|
||||
"""ConvNeXt Block adapted from https://github.com/facebookresearch/ConvNeXt to 1D audio signal."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -48,7 +38,7 @@ class ConvNeXtBlock(nn.Module):
|
||||
dw_kernel_size: int = 7,
|
||||
):
|
||||
super().__init__()
|
||||
self.dwconv = nn.Conv1d(dim, dim, kernel_size=dw_kernel_size, padding=dw_kernel_size//2, groups=dim) # depthwise conv
|
||||
self.dwconv = nn.Conv1d(dim, dim, kernel_size=dw_kernel_size, padding=dw_kernel_size//2, groups=dim)
|
||||
self.norm = nn.LayerNorm(dim, eps=1e-6)
|
||||
self.mlp = SimpleMLP(dim, intermediate_dim)
|
||||
self.gamma = (
|
||||
@@ -72,17 +62,6 @@ class ConvNeXtBlock(nn.Module):
|
||||
|
||||
|
||||
class VocosBackbone(nn.Module):
|
||||
"""
|
||||
Vocos backbone module built with ConvNeXt blocks.
|
||||
|
||||
Args:
|
||||
input_channels (int): Number of input features channels.
|
||||
dim (int): Hidden dimension of the model.
|
||||
intermediate_dim (int): Intermediate dimension used in ConvNeXtBlock.
|
||||
num_layers (int): Number of ConvNeXtBlock layers.
|
||||
layer_scale_init_value (float, optional): Initial value for layer scaling.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_channels: int,
|
||||
@@ -124,16 +103,7 @@ class VocosBackbone(nn.Module):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x (Tensor): Input tensor of shape (B, C, L), where B is the batch size,
|
||||
C denotes output features, and L is the sequence length.
|
||||
|
||||
Returns:
|
||||
Tensor: Output of shape (B, L, H), where B is the batch size, L is the sequence length,
|
||||
and H denotes the model dimension.
|
||||
"""
|
||||
x = self.embed(x) # (B, C, L)
|
||||
x = self.embed(x)
|
||||
x = self.norm(x.transpose(1, 2))
|
||||
x = x.transpose(1, 2)
|
||||
for conv_block in self.convnext:
|
||||
@@ -185,6 +155,7 @@ class Encoder(nn.Module):
|
||||
|
||||
def encode(self, x):
|
||||
x = self.encoder(x)
|
||||
# Sequence slicing for downsampling
|
||||
x = x[:, :, ::self.downsample_scale]
|
||||
x = x.transpose(1, 2)
|
||||
x = self.downsampler(x)
|
||||
@@ -192,10 +163,11 @@ class Encoder(nn.Module):
|
||||
return x
|
||||
|
||||
def preprocess(self, audio):
|
||||
if audio.dim() == 2: # raw audio
|
||||
# Ensure mel_spec is on the same device as the input audio
|
||||
if audio.dim() == 2: # raw audio (B, T)
|
||||
x = self.mel_spec(audio)
|
||||
x = safe_log(x)
|
||||
elif audio.dim() == 3: # mel spectrogram
|
||||
elif audio.dim() == 3: # mel spectrogram (B, C, T)
|
||||
x = audio
|
||||
return x
|
||||
|
||||
|
||||
+17
-13
@@ -1,7 +1,6 @@
|
||||
"""
|
||||
Adapted from https://github.com/duchenzhuang/FSQ-pytorch/blob/main/quantizers/fsq.py#L41
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from einops import rearrange
|
||||
@@ -12,16 +11,24 @@ class FSQSTE(nn.Module):
|
||||
super().__init__()
|
||||
if levels:
|
||||
self.dim = len(levels)
|
||||
self.levels = torch.tensor(levels, dtype=torch.int32).view(1, 1, self.dim)
|
||||
self.half_levels = (self.levels - 1) * (1 - 1e-3) / 2
|
||||
self.offset = 0.5 - 0.5 * (self.levels % 2)
|
||||
self.shift = torch.tan(self.offset / self.half_levels)
|
||||
else:
|
||||
self.levels = levels
|
||||
|
||||
self._basis = torch.cumprod(torch.tensor([1] + levels[:-1]),
|
||||
dim=0,
|
||||
dtype=torch.int32)
|
||||
# Use register_buffer so these move to GPU automatically
|
||||
levels_tensor = torch.tensor(levels, dtype=torch.int32).view(1, 1, self.dim)
|
||||
self.register_buffer("levels", levels_tensor)
|
||||
|
||||
half_levels = (self.levels - 1) * (1 - 1e-3) / 2
|
||||
self.register_buffer("half_levels", half_levels)
|
||||
|
||||
offset = 0.5 - 0.5 * (self.levels % 2)
|
||||
self.register_buffer("offset", offset)
|
||||
|
||||
shift = torch.tan(self.offset / self.half_levels)
|
||||
self.register_buffer("shift", shift)
|
||||
|
||||
basis = torch.cumprod(torch.tensor([1] + levels[:-1]), dim=0, dtype=torch.int32)
|
||||
self.register_buffer("_basis", basis)
|
||||
else:
|
||||
self.levels = None
|
||||
|
||||
def _scale_and_shift(self, zhat_normalized):
|
||||
half_width = self.levels // 2
|
||||
@@ -32,20 +39,17 @@ class FSQSTE(nn.Module):
|
||||
return (zhat - half_width) / half_width
|
||||
|
||||
def indices_to_level_indices(self, indices):
|
||||
""" Converts indices to indices at each level, perhaps needed for a transformer with factorized embeddings """
|
||||
indices = rearrange(indices, '... -> ... 1')
|
||||
codes_non_centered = (indices // self._basis) % self.levels
|
||||
return codes_non_centered
|
||||
|
||||
def to_codebook_index(self, zhat):
|
||||
""" Converts a `code` to an index in the codebook. """
|
||||
assert zhat.shape[-1] == self.dim
|
||||
zhat = self._scale_and_shift(zhat)
|
||||
indices = (zhat * self._basis).sum(dim=-1).round().to(torch.int32)
|
||||
return indices
|
||||
|
||||
def from_codebook_index(self, indices):
|
||||
""" Inverse of `codes_to_indices`. """
|
||||
level_indices = self.indices_to_level_indices(indices)
|
||||
codes = self._scale_and_shift_inverse(level_indices)
|
||||
return codes
|
||||
|
||||
+31
-17
@@ -21,7 +21,7 @@ from huggingface_hub import hf_hub_download
|
||||
|
||||
from encoder.codec import Encoder
|
||||
|
||||
|
||||
# Constants
|
||||
SAMPLE_RATE = 32000
|
||||
SEED = 42
|
||||
VAL_PROP = 0.1
|
||||
@@ -39,35 +39,50 @@ def get_args():
|
||||
def main():
|
||||
args = get_args()
|
||||
input_dir = args.input_dir
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
print("Loading model.")
|
||||
encoder = Encoder()
|
||||
print(f"Loading model onto {device}.")
|
||||
encoder = Encoder().to(device)
|
||||
encoder_path = hf_hub_download(repo_id='ekwek/Soprano-Encoder', filename='encoder.pth')
|
||||
encoder.load_state_dict(torch.load(encoder_path))
|
||||
encoder.load_state_dict(torch.load(encoder_path, map_location=device))
|
||||
encoder.eval()
|
||||
print("Model loaded.")
|
||||
|
||||
|
||||
print("Reading metadata.")
|
||||
files = []
|
||||
with open(f'{input_dir}/metadata.txt', encoding='utf-8') as f:
|
||||
data = f.read().split('\n')
|
||||
metadata_path = input_dir / 'metadata.txt'
|
||||
with open(metadata_path, encoding='utf-8') as f:
|
||||
data = f.read().strip().split('\n')
|
||||
for line in data:
|
||||
if '|' not in line:
|
||||
continue
|
||||
filename, transcript = line.split('|', maxsplit=1)
|
||||
files.append((filename, transcript))
|
||||
print(f'{len(files)} samples located in directory.')
|
||||
|
||||
print("Encoding audio.")
|
||||
print(f"Encoding audio on {device}...")
|
||||
dataset = []
|
||||
for sample in tqdm(files):
|
||||
filename, transcript = sample
|
||||
sr, audio = wavfile.read(f'{input_dir}/wavs/{filename}.wav')
|
||||
audio = torch.from_numpy(audio)
|
||||
for filename, transcript in tqdm(files):
|
||||
wav_path = input_dir / 'wavs' / f'{filename}.wav'
|
||||
|
||||
try:
|
||||
sr, audio = wavfile.read(wav_path)
|
||||
except FileNotFoundError:
|
||||
continue
|
||||
|
||||
# Convert to float tensor for torchaudio/encoder compatibility
|
||||
audio = torch.from_numpy(audio).float()
|
||||
|
||||
if sr != SAMPLE_RATE:
|
||||
audio = torchaudio.functional.resample(audio, sr, SAMPLE_RATE)
|
||||
audio = audio.unsqueeze(0)
|
||||
|
||||
# Prepare for encoder
|
||||
audio = audio.unsqueeze(0).to(device)
|
||||
with torch.no_grad():
|
||||
audio_tokens = encoder(audio)
|
||||
dataset.append([transcript, audio_tokens.squeeze(0).tolist()])
|
||||
|
||||
# Store results (move back to CPU for JSON serialization)
|
||||
dataset.append([transcript, audio_tokens.squeeze(0).cpu().tolist()])
|
||||
|
||||
print("Generating train/test splits.")
|
||||
random.seed(SEED)
|
||||
@@ -79,12 +94,11 @@ def main():
|
||||
print(f'# val samples: {len(val_dataset)}')
|
||||
|
||||
print("Saving datasets.")
|
||||
with open(f'{input_dir}/train.json', 'w', encoding='utf-8') as f:
|
||||
with open(input_dir / 'train.json', 'w', encoding='utf-8') as f:
|
||||
json.dump(train_dataset, f, indent=2)
|
||||
with open(f'{input_dir}/val.json', 'w', encoding='utf-8') as f:
|
||||
with open(input_dir / 'val.json', 'w', encoding='utf-8') as f:
|
||||
json.dump(val_dataset, f, indent=2)
|
||||
print("Datasets saved.")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
[project]
|
||||
name = "soprano-factory"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"transformers",
|
||||
"numpy",
|
||||
"scipy",
|
||||
"tqdm",
|
||||
"einops",
|
||||
"pyyaml",
|
||||
"huggingface-hub",
|
||||
"torch",
|
||||
"torchaudio",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
# Use uv sync --extra gpu to install CUDA support
|
||||
gpu = [
|
||||
"torch",
|
||||
"torchaudio",
|
||||
]
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "pytorch-cu126"
|
||||
url = "https://download.pytorch.org/whl/cu126"
|
||||
explicit = true
|
||||
|
||||
[tool.uv.sources]
|
||||
# Only use the CUDA index when the 'gpu' extra is requested
|
||||
torch = [
|
||||
{ index = "pytorch-cu126", extra = "gpu" },
|
||||
]
|
||||
torchaudio = [
|
||||
{ index = "pytorch-cu126", extra = "gpu" },
|
||||
]
|
||||
@@ -6,3 +6,4 @@ torch
|
||||
torchaudio
|
||||
tqdm
|
||||
transformers
|
||||
pyyaml
|
||||
@@ -14,7 +14,8 @@ import argparse
|
||||
import pathlib
|
||||
import random
|
||||
import time
|
||||
|
||||
import shutil
|
||||
import yaml
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
@@ -23,41 +24,30 @@ from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
from dataset import AudioDataset
|
||||
|
||||
|
||||
def get_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--input-dir",
|
||||
required=False,
|
||||
default="./example_dataset",
|
||||
type=pathlib.Path
|
||||
)
|
||||
parser.add_argument("--save-dir",
|
||||
required=True,
|
||||
type=pathlib.Path
|
||||
)
|
||||
parser.add_argument("--config", type=str, default="config.yaml", help="Path to config file")
|
||||
parser.add_argument("--input-dir", required=False, default="./example_dataset", type=pathlib.Path)
|
||||
parser.add_argument("--save-dir", required=True, type=pathlib.Path)
|
||||
return parser.parse_args()
|
||||
|
||||
# --- Load Configuration ---
|
||||
args = get_args()
|
||||
|
||||
# training hyperparameters
|
||||
device = 'cuda:0'
|
||||
seed = 1337
|
||||
max_lr = 5e-4
|
||||
warmup_ratio = 0.1
|
||||
cooldown_ratio = 0.1
|
||||
with open(args.config, 'r') as f:
|
||||
config = yaml.safe_load(f)
|
||||
|
||||
# Inject YAML keys into global
|
||||
globals().update(config)
|
||||
|
||||
# Handle derived values and type casting
|
||||
betas = tuple(betas) if isinstance(betas, list) else betas
|
||||
min_lr = 0.1 * max_lr
|
||||
batch_size = 4
|
||||
grad_accum_steps = 1
|
||||
seq_len = 1024
|
||||
val_freq = 250
|
||||
text_factor = 0.0 # currently does not train on text inputs, you can increase to change this
|
||||
max_steps = 10000
|
||||
betas = (0.9, 0.95)
|
||||
weight_decay = 0.1
|
||||
train_dataset_path = f'{args.input_dir}/train.json'
|
||||
val_dataset_path = f'{args.input_dir}/val.json'
|
||||
train_dataset_path = args.input_dir / 'train.json'
|
||||
val_dataset_path = args.input_dir / 'val.json'
|
||||
save_path = args.save_dir
|
||||
|
||||
# --- Setup Utilities ---
|
||||
def worker_seed_init(_):
|
||||
worker_seed = torch.initial_seed() % (2**32-1)
|
||||
np.random.seed(worker_seed)
|
||||
@@ -83,10 +73,9 @@ def collate_pack(texts):
|
||||
cur_sample, cur_size = [], 0
|
||||
if len(batch) == batch_size:
|
||||
break
|
||||
if cur_sample and not batch: # add partial sample if there isn't enough data
|
||||
if cur_sample and not batch:
|
||||
batch.append(torch.cat(cur_sample + [torch.zeros(seq_len, dtype=torch.long)])[: seq_len + 1])
|
||||
if len(batch) < batch_size:
|
||||
# pad up to batch_size for consistency
|
||||
pad = batch[-1]
|
||||
while len(batch) < batch_size:
|
||||
batch.append(pad)
|
||||
@@ -108,7 +97,7 @@ def compute_loss(logits, y, num_steps):
|
||||
acc = acc / num_steps
|
||||
return audio_loss, text_loss, acc
|
||||
|
||||
def evaluate(val_dataloader):
|
||||
def evaluate(val_dataloader, model, device, device_type):
|
||||
model.eval()
|
||||
val_dataloader_it = iter(val_dataloader)
|
||||
with torch.no_grad():
|
||||
@@ -128,29 +117,35 @@ def evaluate(val_dataloader):
|
||||
print(f"validation text loss: {val_text_loss_accum.item():.4f}\tvalidation audio loss: {val_audio_loss_accum.item():.4f}\tvalidation acc: {val_acc_accum.item():.4f}")
|
||||
model.train()
|
||||
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained('ekwek/Soprano-80M')
|
||||
# --- Main Training Loop ---
|
||||
if __name__ == '__main__':
|
||||
# Initialize tokenizer
|
||||
tokenizer = AutoTokenizer.from_pretrained('ekwek/Soprano-80M')
|
||||
|
||||
# Environment Setup
|
||||
device_type = "cuda" if device.startswith("cuda") else "cpu"
|
||||
torch.manual_seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.set_float32_matmul_precision('high')
|
||||
|
||||
# Create Save Directory & archive the config used
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
shutil.copy(args.config, save_path / "config_used.yaml")
|
||||
print(f"Save Path: {save_path}")
|
||||
|
||||
# lr schedule
|
||||
# LR schedule steps
|
||||
warmup_steps = int(max_steps * warmup_ratio)
|
||||
cooldown_steps = int(max_steps * cooldown_ratio)
|
||||
|
||||
# model
|
||||
# Model
|
||||
model = AutoModelForCausalLM.from_pretrained('ekwek/Soprano-80M')
|
||||
model.to(torch.bfloat16).to(device)
|
||||
model.train()
|
||||
|
||||
# dataset
|
||||
dataset = AudioDataset(train_dataset_path)
|
||||
# we need batch_size * 16 to have enough tokens after packing
|
||||
dataloader = DataLoader(dataset,
|
||||
# Datasets
|
||||
dataloader = DataLoader(
|
||||
AudioDataset(train_dataset_path),
|
||||
batch_size=batch_size * 16,
|
||||
shuffle=True,
|
||||
num_workers=4,
|
||||
@@ -160,8 +155,9 @@ if __name__ == '__main__':
|
||||
collate_fn=collate_pack,
|
||||
)
|
||||
dataloader_it = iter(dataloader)
|
||||
val_dataset = AudioDataset(val_dataset_path)
|
||||
val_dataloader = DataLoader(val_dataset,
|
||||
|
||||
val_dataloader = DataLoader(
|
||||
AudioDataset(val_dataset_path),
|
||||
batch_size=batch_size * 16,
|
||||
shuffle=False,
|
||||
num_workers=1,
|
||||
@@ -171,33 +167,37 @@ if __name__ == '__main__':
|
||||
collate_fn=collate_pack,
|
||||
)
|
||||
|
||||
# optimizer
|
||||
# Optimizer
|
||||
opt = torch.optim.AdamW(model.parameters(), max_lr, betas=betas, weight_decay=weight_decay, fused=True)
|
||||
|
||||
pbar = tqdm(range(0, max_steps), ncols=200, dynamic_ncols=True)
|
||||
for step in pbar:
|
||||
start = time.time()
|
||||
if val_freq > 0 and (step % val_freq == 0 or step == max_steps - 1):
|
||||
evaluate(val_dataloader)
|
||||
evaluate(val_dataloader, model, device, device_type)
|
||||
|
||||
opt.zero_grad()
|
||||
audio_loss_accum = 0.0
|
||||
text_loss_accum = 0.0
|
||||
acc_accum = 0.0
|
||||
|
||||
for micro_step in range(grad_accum_steps):
|
||||
try:
|
||||
x, y = next(dataloader_it)
|
||||
except:
|
||||
except StopIteration:
|
||||
dataloader_it = iter(dataloader)
|
||||
x, y = next(dataloader_it)
|
||||
|
||||
x, y = x.to(device), y.to(device)
|
||||
|
||||
with torch.autocast(device_type=device_type, dtype=torch.bfloat16):
|
||||
logits = model(x).logits
|
||||
audio_loss, text_loss, acc = compute_loss(logits, y, grad_accum_steps)
|
||||
|
||||
audio_loss_accum += audio_loss.detach()
|
||||
text_loss_accum += text_loss.detach()
|
||||
acc_accum += acc.detach()
|
||||
|
||||
total_loss = audio_loss + text_factor * text_loss
|
||||
total_loss.backward()
|
||||
|
||||
@@ -205,16 +205,20 @@ if __name__ == '__main__':
|
||||
lr = get_lr(step)
|
||||
for param_group in opt.param_groups:
|
||||
param_group['lr'] = lr
|
||||
|
||||
opt.step()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
total_tokens = step * batch_size*seq_len*grad_accum_steps
|
||||
|
||||
end = time.time()
|
||||
dt = (end - start) * 1000
|
||||
tokens_per_second = (batch_size * seq_len * grad_accum_steps) / (end - start)
|
||||
tqdm_log = f'text loss: {text_loss_accum.item():.3f} | audio loss: {audio_loss_accum.item():.3f} | acc: {acc_accum.item():.4f} | lr: {lr:.2e} | norm: {norm:.3f} | time: {dt:.2f} ms | {tokens_per_second:.2f} t/s'
|
||||
|
||||
tqdm_log = f'text: {text_loss_accum.item():.3f} | audio: {audio_loss_accum.item():.3f} | acc: {acc_accum.item():.4f} | lr: {lr:.2e} | ms: {dt:.1f} | t/s: {tokens_per_second:.1f}'
|
||||
pbar.set_description(tqdm_log)
|
||||
|
||||
print(f"Training complete. Saving model at {save_path}")
|
||||
model.save_pretrained(save_path)
|
||||
tokenizer.save_pretrained(save_path)
|
||||
print("Saving done.")
|
||||
|
||||
Reference in New Issue
Block a user