如何避免WSL训练Wave U-Net模型时出现冻结崩溃问题?
WSL训练Wave U-Net时冻结崩溃的解决方法
问题描述
在WSL环境下训练基于MSE损失的Wave U-Net神经网络时遇到以下问题:
- 最初通过VS Code的Jupyter内核逐行运行脚本,训练10个epoch后,WSL崩溃并陷入无限重连循环;
- 改用终端直接运行脚本后,进程卡顿超过1小时。
模型代码
import os import torch import torch.nn as nn import torch.nn.init as init import torch.nn.functional as F import torch.optim as optim import torch.nn as nn from torch.utils.data import Dataset, DataLoader,random_split import torchaudio import pandas as pd from sklearn.preprocessing import MinMaxScaler import numpy as np import matplotlib.pyplot as plt PARENT_FOLDER = "/mnt/c/Users/Tudor/Documents/yt-dlp" SCALER = MinMaxScaler() class DownSamplingLayer(nn.Module): def __init__(self, channel_in, channel_out, dilation=1, kernel_size=9, stride=1, padding="same"): super(DownSamplingLayer, self).__init__() self.main = nn.Sequential( nn.Conv1d(channel_in, channel_out, kernel_size=kernel_size, stride=stride, padding=padding, dilation=dilation), nn.BatchNorm1d(channel_out), nn.LeakyReLU(negative_slope=0.1, inplace=True), ) self.dropout = nn.Dropout(p=0.3) def forward(self, x): x = self.main(x) return self.dropout(x) class UpSamplingLayer(nn.Module): def __init__(self, channel_in, channel_out, kernel_size=9, stride=1, padding="same"): super(UpSamplingLayer, self).__init__() self.main = nn.Sequential( nn.Conv1d(channel_in, channel_out, kernel_size=kernel_size, stride=stride, padding=padding), nn.BatchNorm1d(channel_out), nn.LeakyReLU(negative_slope=0.1, inplace=True), ) def forward(self, x): return self.main(x) class Model(nn.Module): def __init__(self, n_layers=8, channels_interval=16): super(Model, self).__init__() self.n_layers = n_layers self.channels_interval = channels_interval encoder_in_channels_list = [1] + [i * self.channels_interval for i in range(1, self.n_layers)] encoder_out_channels_list = [i * self.channels_interval for i in range(1, self.n_layers + 1)] self.encoder = nn.ModuleList() for i in range(self.n_layers): self.encoder.append( DownSamplingLayer( channel_in=encoder_in_channels_list[i], channel_out=encoder_out_channels_list[i] ) ) self.middle = nn.Sequential( nn.Conv1d(self.n_layers * self.channels_interval, self.n_layers * self.channels_interval, kernel_size=3, stride=1, padding="same"), nn.BatchNorm1d(self.n_layers * self.channels_interval), nn.LeakyReLU(negative_slope=0.1, inplace=True) ) decoder_in_channels_list = [(2 * i + 1) * self.channels_interval for i in range(1, self.n_layers)] + [ 2 * self.n_layers * self.channels_interval] decoder_in_channels_list = decoder_in_channels_list[::-1] decoder_out_channels_list = encoder_out_channels_list[::-1] self.decoder = nn.ModuleList() for i in range(self.n_layers): self.decoder.append( UpSamplingLayer( channel_in=decoder_in_channels_list[i], channel_out=decoder_out_channels_list[i] ) ) self.out = nn.Sequential( nn.Conv1d(1+self.channels_interval, 1, kernel_size=1, stride=1), nn.LeakyReLU(negative_slope=0.1, inplace=True) ) # Initialize the weights self.initialize_weights() def initialize_weights(self): for m in self.modules(): if isinstance(m, nn.Conv1d) or isinstance(m, nn.ConvTranspose1d): init.xavier_uniform_(m.weight, gain=1.0) if m.bias is not None: init.constant_(m.bias, 0) elif isinstance(m, nn.BatchNorm1d): init.constant_(m.weight, 1) init.constant_(m.bias, 0) def forward(self, x): tmp = [] o = x # Up Sample for i in range(self.n_layers): o = self.encoder[i](o) tmp.append(o) o = F.max_pool1d(o, kernel_size=2, stride=2) o = self.middle(o) for i in range(self.n_layers): o = F.interpolate(o, scale_factor=2, mode="linear", align_corners=True) o = torch.cat((o, tmp[self.n_layers - i - 1]), dim=1) o = self.decoder[i](o) o = torch.cat((o, x), dim=1) o = self.out(o) return o input_samples, target_samples = np.load("input_samples.npy"), np.load("target_samples.npy") input_samples =input_samples.tolist() target_samples =target_samples.tolist() df = pd.DataFrame({"input":input_samples,"target":target_samples}) class BassenhanceDataset(Dataset): def __init__(self, df): self.df = df self.input = df["input"] self.target = df["target"] def __len__(self): return len(self.df) def __getitem__(self, idx): input = self.input[idx] target = self.target[idx] input = torch.tensor(input, dtype=torch.float32).T target = torch.tensor(target, dtype=torch.float32).T return input, target def get_loader(self, batch_size, shuffle=True, num_workers=0): return DataLoader(self, batch_size=batch_size, shuffle=shuffle, num_workers=num_workers, collate_fn=self.collate_fn) def transpose(self, data): return data.transpose(1,2) def get_loader_transpose(self, batch_size, shuffle=True, num_workers=0): return DataLoader(self, batch_size=batch_size, shuffle=shuffle, num_workers=num_workers, collate_fn=self.collate_fn, drop_last=True) def split(self, train_size=0.8, shuffle=True): return torch.utils.data.random_split(self, [int(len(self) * train_size), len(self) - int(len(self) * train_size)], generator=torch.Generator().manual_seed(42)) def train_epoch(model, train_loader, optimizer, criterion, device): model.train() running_loss = 0.0 for i, (input, target) in enumerate(train_loader): input, target = input.to(device), target.to(device) optimizer.zero_grad() output = model(input) loss = criterion(output, target) loss.backward() optimizer.step() running_loss += loss.item() return running_loss / len(train_loader) def validate_epoch(model, val_loader, criterion, device): model.eval() running_loss = 0.0 with torch.no_grad(): for i, (input, target) in enumerate(val_loader): input, target = input.to(device), target.to(device) output = model(input) loss = criterion(output, target) running_loss += loss.item() return running_loss / len(val_loader) def train(model, train_loader, val_loader, optimizer, criterion, device, epochs=10): train_losses = [] val_losses = [] for epoch in range(epochs): train_loss = train_epoch(model, train_loader, optimizer, criterion, device) val_loss = validate_epoch(model, val_loader, criterion, device) train_losses.append(train_loss) val_losses.append(val_loss) print(f"Epoch {epoch + 1} | Train Loss: {train_loss:.10f} | Val Loss: {val_loss:.10f}") save_state(model, epoch + 1) if early_stopping(val_losses, patience=50): print("Early Stopping") break return train_losses, val_losses def plot_losses_real_time(train_losses, val_losses): plt.plot(train_losses, label="Train Loss") plt.plot(val_losses, label="Val Loss") plt.title("Losses") plt.xlabel("Epoch") plt.ylabel("Loss") plt.legend() plt.show() def early_stopping(val_losses, patience=5): if len(val_losses) < patience: return False else: return val_losses[-1] > val_losses[-2] > val_losses[-3] def save_state(model, epoch, path = "models"): if epoch % 10 == 0: state = { "epoch": epoch, "state_dict": model.state_dict(), "optimizer": optimizer.state_dict() } torch.save(state, os.path.join(path, f"model_{epoch}.pth")) print("Saved model") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Using {device} device") model = Model(5,32).to(device) # Mean Squared Error Loss criterion = nn.MSELoss() optimizer = optim.Adam(model.parameters(), lr=0.00002) dataset = BassenhanceDataset(df) train_dataset, val_dataset = dataset.split() print(f"Train dataset size: {len(train_dataset)}") print(f"Val dataset size: {len(val_dataset)}") train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=2) valid_loader = DataLoader(val_dataset, batch_size=128, shuffle=True, num_workers=2) train_losses, val_losses = train(model, train_loader, valid_loader, optimizer, criterion, device, epochs=100) plot_losses_real_time(train_losses, val_losses) save_state(model, 100)
解决方法
一、优化WSL系统配置
- 限制WSL内存占用:在Windows用户目录下创建
.wslconfig文件,添加以下内容后重启WSL(执行wsl --shutdown再重新打开):[wsl2] memory=8GB swap=4GB localhostForwarding=true - 关闭实时扫描或添加排除项:临时关闭Windows Defender实时保护,或者把WSL数据目录加入扫描排除列表,减少IO拖慢。
- 切换到WSL2:执行
wsl -l -v查看版本,若为WSL1,运行wsl --set-version <发行版名称> 2切换,WSL2的IO和性能远优于WSL1。
二、调整训练脚本
- 降低batch size:当前batch size=128,若显存不足会导致频繁内存交换,先降到32或16测试。
- 优化DataLoader参数:
- 将
num_workers改为0,WSL下多进程数据加载容易出兼容性问题; - 启用
pin_memory=True(仅CUDA环境),减少CPU到GPU的复制开销; - 添加
persistent_workers=True,避免每个epoch重建数据加载进程。
- 将
- 修复代码缺陷:
save_state函数直接引用全局optimizer,改为传入参数:
调用时改为def save_state(model, optimizer, epoch, path = "models"): os.makedirs(path, exist_ok=True) if epoch % 10 == 0: state = { "epoch": epoch, "state_dict": model.state_dict(), "optimizer": optimizer.state_dict() } torch.save(state, os.path.join(path, f"model_{epoch}.pth")) print("Saved model")save_state(model, optimizer, epoch + 1);- 修正
early_stopping逻辑,改为判断连续patience个epoch损失无下降:def early_stopping(val_losses, patience=5, min_delta=0.0): if len(val_losses) < patience: return False recent_losses = val_losses[-patience:] return all(recent_losses[i] <= recent_losses[i+1] + min_delta for i in range(patience-1))
- 替换实时绘图:WSL下
plt.show()易阻塞进程,改为保存图片到文件:def plot_losses(train_losses, val_losses): plt.plot(train_losses, label="Train Loss") plt.plot(val_losses, label="Val Loss") plt.title("Losses") plt.xlabel("Epoch") plt.ylabel("Loss") plt.legend() plt.savefig("loss_curve.png") plt.close()
三、优化训练环境
- 迁移数据集到WSL本地:把
/mnt/c/...下的数据集复制到WSL本地目录(如~/data/),避免跨系统IO的性能损耗。 - 后台运行脚本:用
nohup让脚本在后台执行,避免终端卡顿影响:
可通过nohup python train.py > train.log 2>&1 &tail -f train.log实时查看训练日志。 - 监控资源状态:用
htop查看CPU、内存占用,用nvidia-smi查看GPU状态,确认是否有资源耗尽情况。
内容的提问来源于stack exchange,提问作者user22818756
相关产品推荐
相关产品推荐

