You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何避免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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.07 16:15:54