Mô hình Transformer đã trở nên phổ biến kể từ năm 2017 và đã được ứng dụng rộng rãi trong nhiều lĩnh vực. Trong bài viết này, chúng ta sẽ phân tích ưu nhược điểm của Transformer trước khi chuyển sang mô hình Manba mới.
Mô hình Transformer có cấu trúc như Hình 1, nơi mà cơ chế chú ý tự là thành phần quan trọng nhất. Cơ chế này chia thành các thành phần Q, K và V, trong đó V là câu trả lời cho Q. Thường thì K và Q giống nhau, nhưng cũng có thể khác nhau, miễn là K cần có liên quan đến V. Q và K tính toán mức độ tương quan giữa chúng, tạo ra một ma trận hai chiều lớn, biểu thị tương quan giữa Qi và Kj. Sau đó, ma trận tương quan này được nhân với V để tạo ra giá trị mới cho Vj dựa trên các giá trị khác, tức là Vj = v1 * V1 + v2 * V2 + .... Điều này giúp mô hình tập trung vào mối quan hệ ngữ cảnh hơn, từ đó dự đoán token tiếp theo một cách chính xác hơn. Tuy nhiên, mô hình Transformer cũng có những hạn chế, đặc biệt là độ phức tạp tính toán cao, đạt O(N^2 * d + N^2), với d là kích thước nhúng và N là độ dài ngữ cảnh. Mặc dù có thể tăng tốc quá trình huấn luyện bằng cách tính toán song song, nhưng quá trình suy luận lại chậm. Vì vậy, các nhà khoa học đang tìm kiếm giải pháp cải tiến.
Năm 2023, một mô hình mới được giới thiệu gọi là mô hình Manba. Mô hình này không chỉ nhanh hơn trong việc huấn luyện mà còn hiệu quả trong việc suy luận, đồng thời cải thiện khả năng xử lý ngữ cảnh dài hơn so với Transformer.
Mô Hình SSM (Structured State Space Model)
Trước khi đi sâu vào mô hình Manba, chúng ta cần hiểu về mô hình SSM. Đây là mô hình liên quan trực tiếp đến Manba. Công thức cơ bản của mô hình SSM như sau:
A đại diện cho ma trận chuyển trạng thái. Chúng ta coi f(x) = tổng (cn * e(i2pi * n)) là một hàm đường cong trong không gian nhiều chiều. Khi chuyển đổi thành dạng rời rạc, công thức sẽ thay đổi thành A = e(delta * A) và B sẽ được tính toán dựa trên phương pháp bảo trì bậc 0.
Tuy nhiên, nếu A, B, C không thay đổi, mô hình trở thành một hệ thống tuyến tính không biến đổi, không có khả năng tập trung vào các thông tin cụ thể. Để khắc phục điều này, mô hình Manba đưa vào hệ thống tuyến tính biến đổi, làm cho B, C và delta phụ thuộc vào đầu vào, từ đó có thể lọc hoặc tập trung vào các thông tin cụ thể.
Xem Nguyên Lý Mô Hình Manba Qua Mã Nguồn
Hình 5: Cấu trúc mô hình Manba
Hình 6: Mamba Block
def tao_block(
kich_thuoc_mo_hinh,
kich_thuoc_trung_gian,
cau_hinh_ssm=None,
chi_so_attn=None,
cau_hinh_attn=None,
epsilon_luong_chuan=1e-5,
dung_luong_RMS=False,
du_lieu_cong_thuc_cong=False,
hop_nhat_cong_chuan=False,
chi_so_layer=None,
thiet_bi=None,
kieu_du_lieu=None,):
if cau_hinh_ssm is None:
cau_hinh_ssm = {}
if chi_so_attn is None:
chi_so_attn = []
if cau_hinh_attn is None:
cau_hinh_attn = {}
tham_so_tao_ra = {"thiet_bi": thiet_bi, "kieu_du_lieu": kieu_du_lieu}
if chi_so_layer not in chi_so_attn:
cau_hinh_ssm = copy.deepcopy(cau_hinh_ssm) if cau_hinh_ssm is not None else {}
loai_ssm = cau_hinh_ssm.pop("loai", "Mamba1")
if loai_ssm not in ["Mamba1", "Mamba2"]:
raise ValueError(f"Loại ssm không hợp lệ: {loai_ssm}, chỉ hỗ trợ Mamba1 và Mamba2")
lop_phan_chia = partial(
Mamba2 if loai_ssm == "Mamba2" else Mamba,
chi_so_layer=chi_so_layer,
**cau_hinh_ssm,
**tham_so_tao_ra
)
else:
lop_phan_chia = partial(MHA, chi_so_layer=chi_so_layer, **cau_hinh_attn, **tham_so_tao_ra)
lop_chuan_hoa = partial(
nn.LayerNorm if not dung_luong_RMS else RMSNorm, eps=epsilon_luong_chuan, **tham_so_tao_ra
)
if kich_thuoc_trung_gian == 0:
lop_mlp = nn.Identity
else:
lop_mlp = partial(
GatedMLP, so_tinh_trung_gian=kich_thuoc_trung_gian, so_tinh_ra=kich_thuoc_mo_hinh, **tham_so_tao_ra
)
khoi = Khoi(
kich_thuoc_mo_hinh,
lop_phan_chia,
lop_mlp,
lop_chuan_hoa=lop_chuan_hoa,
hop_nhat_cong_chuan=hop_nhat_cong_chuan,
du_lieu_cong_thuc_cong=du_lieu_cong_thuc_cong, )
khoi.chi_so_layer = chi_so_layer
return khoi
class Khoi(nn.Module):
def __init__(
self, kich_thuoc, lop_phan_chia, lop_mlp, lop_chuan_hoa=nn.LayerNorm, hop_nhat_cong_chuan=False, du_lieu_cong_thuc_cong=False ):
super().__init__()
self.du_lieu_cong_thuc_cong = du_lieu_cong_thuc_cong
self.hop_nhat_cong_chuan = hop_nhat_cong_chuan
self.chuan_hoa = lop_chuan_hoa(kich_thuoc)
self.phan_chia = lop_phan_chia(kich_thuoc)
if lop_mlp is not nn.Identity:
self.chuan_hoa2 = lop_chuan_hoa(kich_thuoc)
self.mlp = lop_mlp(kich_thuoc)
else:
self.mlp = None
if self.hop_nhat_cong_chuan:
assert RMSNorm is not None, "Không thể nhập RMSNorm"
assert isinstance(
self.chuan_hoa, (nn.LayerNorm, RMSNorm)
), "Chỉ hỗ trợ LayerNorm và RMSNorm cho hop_nhat_cong_chuan"
def forward(
self, trang_thai_an_toan: Tensor, du_lieu_cong_thuc_cong: Optional[Tensor] = None, tham_so_ket_noi=None, **tham_so_phan_chia
):
if not self.hop_nhat_cong_chuan:
du_lieu_cong_thuc_cong = (trang_thai_an_toan + du_lieu_cong_thuc_cong) if du_lieu_cong_thuc_cong is not None else trang_thai_an_toan
trang_thai_an_toan = self.chuan_hoa(du_lieu_cong_thuc_cong.to(dtype=self.chuan_hoa.weight.dtype))
if self.du_lieu_cong_thuc_cong:
du_lieu_cong_thuc_cong = du_lieu_cong_thuc_cong.to(torch.float32)
else:
trang_thai_an_toan, du_lieu_cong_thuc_cong = ham_chuan_hoa(
trang_thai_an_toan,
self.chuan_hoa.weight,
self.chuan_hoa.bias,
du_lieu_cong_thuc_cong=du_lieu_cong_thuc_cong,
prenorm=True,
du_lieu_cong_thuc_cong=self.du_lieu_cong_thuc_cong,
eps=self.chuan_hoa.eps,
la_RMSNorm=isinstance(self.chuan_hoa, RMSNorm)
)
trang_thai_an_toan = self.phan_chia(trang_thai_an_toan, tham_so_ket_noi=tham_so_ket_noi, **tham_so_phan_chia)
if self.mlp is not None:
if not self.hop_nhat_cong_chuan:
du_lieu_cong_thuc_cong = trang_thai_an_toan + du_lieu_cong_thuc_cong
trang_thai_an_toan = self.chuan_hoa2(du_lieu_cong_thuc_cong.to(dtype=self.chuan_hoa2.weight.dtype))
if self.du_lieu_cong_thuc_cong:
du_lieu_cong_thuc_cong = du_lieu_cong_thuc_cong.to(torch.float32)
else:
trang_thai_an_toan, du_lieu_cong_thuc_cong = ham_chuan_hoa(
trang_thai_an_toan,
self.chuan_hoa2.weight,
self.chuan_hoa2.bias,
du_lieu_cong_thuc_cong=du_lieu_cong_thuc_cong,
prenorm=True,
du_lieu_cong_thuc_cong=self.du_lieu_cong_thuc_cong,
eps=self.chuan_hoa2.eps,
la_RMSNorm=isinstance(self.chuan_hoa2, RMSNorm)
)
trang_thai_an_toan = self.mlp(trang_thai_an_toan)
return trang_thai_an_toan, du_lieu_cong_thuc_cong
def forward(self, trang_thai_an_toan, tham_so_ket_noi=None):
batch, do_dai, kich_thuoc = trang_thai_an_toan.shape
trang_thai_conv, trang_thai_ssm = None, None
if tham_so_ket_noi is not None:
trang_thai_conv, trang_thai_ssm = self._lay_trang_thai_tu_bộ_nhớ(tham_so_ket_noi, batch)
if tham_so_ket_noi.do_dai_viet > 0:
ket_qua, _, _ = self.buoc(trang_thai_an_toan, trang_thai_conv, trang_thai_ssm)
return ket_qua
xz = rearrange(
self.in_proj.weight @ rearrange(trang_thai_an_toan, "b l d -> d (b l)"),
"d (b l) -> b d l",
l=do_dai, )
if self.in_proj.bias is not None:
xz = xz + rearrange(self.in_proj.bias.to(dtype=xz.dtype), "d -> d 1")
A = -torch.exp(self.A_log.float())
if self.dung_duong_linh_tinh and ham_conv1d is not None and tham_so_ket_noi is None:
ket_qua = ham_tien_hanh_mamba(
xz,
self.conv1d.weight,
self.conv1d.bias,
self.x_proj.weight,
self.dt_proj.weight,
self.out_proj.weight,
self.out_proj.bias,
A,
None,
None,
self.D.float(),
delta_bias=self.dt_proj.bias.float(),
delta_softplus=True,
)
else:
x, z = xz.chunk(2, dim=1)
if conv_state is not None:
conv_state.copy_(F.pad(x, (self.d_conv - x.shape[-1], 0)))
if ham_conv1d is None:
x = self.ham_kich_hoat(self.conv1d(x)[..., :do_dai])
else:
assert self.ham_kich_hoat in ["silu", "swish"]
x = ham_conv1d(
x=x,
weight=rearrange(self.conv1d.weight, "d 1 w -> d w"),
bias=self.conv1d.bias,
ham_kich_hoat=self.ham_kich_hoat,
)
x_dbl = self.x_proj(rearrange(x, "b d l -> (b l) d"))
dt, B, C = torch.split(x_dbl, [self.dt_rank, self.d_state, self.d_state], dim=-1)
dt = self.dt_proj.weight @ dt.t()
dt = rearrange(dt, "d (b l) -> b d l", l=do_dai)
B = rearrange(B, "(b l) dstate -> b dstate l", l=do_dai).contiguous()
C = rearrange(C, "(b l) dstate -> b dstate l", l=do_dai).contiguous()
assert self.ham_kich_hoat in ["silu", "swish"]
y = ham_quet_chon(
x,
dt,
A,
B,
C,
self.D.float(),
z=z,
delta_bias=self.dt_proj.bias.float(),
delta_softplus=True,
tra_ve_trang_thai_cuoi_cung=trang_thai_ssm is not None,
)
if trang_thai_ssm is not None:
y, trang_thai_cuoi_cung = y
trang_thai_ssm.copy_(trang_thai_cuoi_cung)
y = rearrange(y, "b d l -> b l d")
ket_qua = self.out_proj(y)
return ket_qua
Theo mã nguồn trên và Hình Mamba, ta thấy dữ liệu tại thời điểm x(t) sau khi qua phép biến đổi tuyến tính trở thành Bt, deltaT, Ct. Điều này khiến B, delta, C trở thành các biến phụ thuộc vào đầu vào, trong khi phép biến đổi vẫn giữ nguyên.
Ma trận chuyển trạng thái thay đổi theo delta t, tạo thành một hệ thống biến đổi theo thời gian. A được coi là phần không đổi của hệ thống, đại diện cho logic và phản xạ cố định, trong khi delta kiểm soát phạm vi nhớ hay quên của hệ thống tại mỗi thời điểm.
Phần SSM quan trọng nhất là hàm ham_quet_chon.
def ham_quet_chon_ref(u, delta, A, B, C, D=None, z=None, delta_bias=None, delta_softplus=False, tra_ve_trang_thai_cuoi_cung=False):
dtype_in = u.dtype
u = u.float()
delta = delta.float()
if delta_bias is not None:
delta = delta + delta_bias[..., None].float()
if delta_softplus:
delta = F.softplus(delta)
batch, kich_thuoc, kich_thuoc_trang_thai = u.shape[0], A.shape[0], A.shape[1]
co_B_bien = B.dim() >= 3
co_C_bien = C.dim() >= 3
if A.is_complex():
if co_B_bien:
B = torch.view_as_complex(rearrange(B.float(), "... (L hai) -> ... L hai", hai=2))
if co_C_bien:
C = torch.view_as_complex(rearrange(C.float(), "... (L hai) -> ... L hai", hai=2))
else:
B = B.float()
C = C.float()
x = A.new_zeros((batch, kich_thuoc, kich_thuoc_trang_thai))
ket_qua = []
deltaA = torch.exp(torch.einsum('bkl,kn->bklm', delta, A))
if not co_B_bien:
deltaB_u = torch.einsum('bkl,kn,bkl->bklm', delta, B, u)
else:
if B.dim() == 3:
deltaB_u = torch.einsum('bkl,bnl,bkl->bklm', delta, B, u)
else:
B = repeat(B, "B G N L -> B (G H) N L", H=kich_thuoc // B.shape[1])
deltaB_u = torch.einsum('bkl,bknl,bkl->bklm', delta, B, u)
if co_C_bien and C.dim() == 4:
C = repeat(C, "B G N L -> B (G H) N L", H=kich_thuoc // C.shape[1])
trang_thai_cuoi_cung = None
for i in range(u.shape[2]):
x = deltaA[:, :, i] * x + deltaB_u[:, :, i]
if not co_C_bien:
y = torch.einsum('bkn,kn->bk', x, C)
else:
if C.dim() == 3:
y = torch.einsum('bkn,bn->bk', x, C[:, :, i])
else:
y = torch.einsum('bkn,bkn->bk', x, C[:, :, :, i])
if i == u.shape[2] - 1:
trang_thai_cuoi_cung = x
if y.is_complex():
y = y.real * 2
ket_qua.append(y)
y = torch.stack(ket_qua, dim=2)
out = y if D is None else y + u * rearrange(D, "k -> k 1")
if z is not None:
out = out * F.silu(z)
out = out.to(dtype=dtype_in)
return out if not tra_ve_trang_thai_cuoi_cung else (out, trang_thai_cuoi_cung)
B, delta, C đều phụ thuộc vào thời điểm t, tương ứng với vector thứ t trong mô hình ngôn ngữ sau khi qua phép biến đổi tuyến tính. Delta đóng vai trò như cổng kiểm soát, điều chỉnh mức độ nhớ hay quên của hệ thống và mức độ chú ý vào thông tin đầu vào hiện tại. B quyết định liệu thông tin đầu vào có được đưa vào không gian trạng thái hay không, trong khi C kiểm soát việc thông tin có được đưa ra đầu ra hay không.