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