GPU…GH200を使う

アライナーにMAMBAを入れる→UnetDenoiserの関数にMambaを入れる

Mamba-Unetという先行研究があった

https://arxiv.org/pdf/2402.05079

https://github.com/ziyangwang007/Mamba-UNet/tree/main

動作確認はマルチモーダル基盤トレーニングコード?

class UNetDenoiser(nn.Module):
    def __init__(self, in_channels : int, mid_channels : int, n_layers : int, norm='layer', act= 'SiLU', ):
        super().__init__()
        out_channels = in_channels
        self.down = nn.ModuleList()
        for i in range(n_layers):
            if i == (n_layers - 1):
                self.down.append(ResidualLinear(in_channels, mid_channels, norm=norm, act=act))
            else:
                self.down.append(ResidualLinear(in_channels, in_channels, norm=norm, act=act))

        self.mid = nn.ModuleList()
        for i in range(n_layers):
            self.mid.append(ResidualLinear(mid_channels, mid_channels, norm=norm, act=act))
        
        self.up = nn.ModuleList()
        for i in range(n_layers):
            if i == 0:
                self.up.append(ResidualLinear(mid_channels * 2, out_channels, norm='none', act='Identity'))
            else:
                self.up.append(ResidualLinear(out_channels * 2, out_channels, norm=norm, act=act))
        
    def forward(self, x):
        down_res = []
        for down_layer in self.down:
            x = down_layer(x)
            down_res.append(x)
        
        for mid_layer in self.mid:
            x = mid_layer(x)
        
        down_res.reverse()
        for up_layer, res in zip(self.up, down_res):
            x = up_layer(torch.cat([x, res], dim=-1))
        return x

mid 部分を Mamba に置き換え → 時系列処理に強い

import torch
import torch.nn as nn
from mamba import Mamba2Simple  # Mamba2Simple を適切にインポート

class UNetDenoiserMamba(nn.Module):
    def __init__(self, in_channels: int, mid_channels: int, n_layers: int, norm='layer', act='SiLU'):
        super().__init__()
        out_channels = in_channels
        self.down = nn.ModuleList()
        
        for i in range(n_layers):
            if i == (n_layers - 1):
                self.down.append(ResidualLinear(in_channels, mid_channels, norm=norm, act=act))
            else:
                self.down.append(ResidualLinear(in_channels, in_channels, norm=norm, act=act))

        # Mamba を mid に組み込む
        self.mamba = Mamba2Simple(
            inp_size=mid_channels, 
            size=mid_channels, 
            norm=False, 
            act='Tanh', 
            update_bias=-1
        )

        self.up = nn.ModuleList()
        for i in range(n_layers):
            if i == 0:
                self.up.append(ResidualLinear(mid_channels * 2, out_channels, norm='none', act='Identity'))
            else:
                self.up.append(ResidualLinear(out_channels * 2, out_channels, norm=norm, act=act))

    def forward(self, x):
        down_res = []
        for down_layer in self.down:
            x = down_layer(x)
            down_res.append(x)

        # Mamba で時系列処理
        x = self.mamba(x)

        down_res.reverse()
        for up_layer, res in zip(self.up, down_res):
            x = up_layer(torch.cat([x, res], dim=-1))
        return x