| from ctypes import Union |
| import math, copy, warnings |
| from re import M |
| import torch |
| import numpy as np |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from .low_rank_tensors import ( |
| TensorTrain, |
| TensorTrainMatrix, |
| BatchTensorTrain, |
| GCPTensorTrain, |
| ) |
| from .utils import config_class |
| from .emb_utils import get_cum_prod, tensorized_lookup |
| from typing import List |
|
|
|
|
| ACT2FN = { |
| "gelu": nn.GELU, |
| "tanh": nn.Tanh, |
| "relu": nn.ReLU, |
| "silu": nn.SiLU, |
| } |
|
|
|
|
| def build_embedding_layer(in_features, out_features, config_embedding): |
| if config_embedding.tensorized: |
| return TensorizedEmbedding(in_features, out_features, config_embedding) |
| else: |
| return nn.Embedding(in_features, out_features) |
|
|
|
|
| def build_head_layer(in_features, out_features, config_head): |
| if config_head.tensorized: |
| return TensorizedLinear(in_features, out_features, config_head, bias=False) |
| else: |
| return nn.Linear(in_features, out_features) |
|
|
|
|
| def build_linear_layer(in_features, out_features, config_linear, bias=True): |
| if config_linear.tensorized: |
| return TensorizedLinear(in_features, out_features, config_linear, bias) |
| else: |
| return nn.Linear(in_features, out_features, bias) |
|
|
|
|
| def build_gcp_linear_layer( |
| in_features, out_features, config_linear, bias=True, gcp_idx=None |
| ): |
| if config_linear.tensorized: |
| config_linear_gcp = copy.copy(config_linear) |
| assert gcp_idx is not None |
| setattr(config_linear_gcp, "gcp_idx", gcp_idx) |
| return TensorizedLinear(in_features, out_features, config_linear_gcp, bias) |
| else: |
| warnings.warn( |
| "You are trying to build GCP layers w/o tensorization, degenerate to regular linear with BP" |
| ) |
| return nn.Linear(in_features, out_features, bias) |
|
|
|
|
| class TensorizedEmbedding(nn.Module): |
| def __init__(self, in_features, out_features, config): |
| super(TensorizedEmbedding, self).__init__() |
|
|
| self.in_features = in_features |
| self.out_features = out_features |
| self.shape = config.shape |
| self.factorization = config.factorization |
|
|
| target_stddev = 1.0 |
| config_tensor = config_class( |
| shape=config.shape, ranks=config.ranks, target_sdv=target_stddev |
| ) |
|
|
| if config.factorization == "ttm": |
| if len(self.shape[0]) % 2: |
| raise ValueError(f"Only support even number of dim factorization") |
|
|
| self.tensor = TensorTrainMatrix(config_tensor) |
| self.register_buffer("ind2coord", self.ttm_get_indices_dict()) |
| self.dict_tensor = torch.zeros(in_features, dtype=torch.long) |
| self.dict_range = torch.arange(0, in_features, dtype=torch.long) |
|
|
| elif config.factorization == "tt": |
| self.tensor = TensorTrain(config_tensor) |
|
|
| else: |
| raise ValueError( |
| f"Got {config.factorization}, only support ttm and tt tensorized embedding for now." |
| ) |
|
|
| def set_scale_factors(self, scale_w=1.0): |
| self.scales = torch.nn.ParameterList() |
| self.scale_factors = torch.nn.ParameterList() |
|
|
| if not isinstance(scale_w, list): |
| scale_w = [scale_w] * len(self.tensor.factors) |
| for s in scale_w: |
| self.scale_factors.append( |
| torch.nn.Parameter(torch.tensor(s, requires_grad=True)) |
| ) |
|
|
| self.scale_row = torch.nn.Embedding(self.in_features, 1) |
| self.scale_row.weight.data[:] = 1.0 |
|
|
| def ttm_get_indices_dict(self): |
| m = len(self.shape[0]) // 2 |
| shape = [ |
| np.prod(self.shape[0][:m]), |
| np.prod(self.shape[0][m:]), |
| ] |
| ind2coord = ( |
| torch.tensor( |
| np.array(np.unravel_index(np.arange(self.in_features), shape)), |
| dtype=torch.long, |
| ) |
| .t() |
| .contiguous() |
| ) |
|
|
| return ind2coord |
|
|
| def tt_select_ind(self, factors, inds): |
| m = len(factors) // 2 |
| input_shape = [U.shape[1] for U in factors[:m]] |
| output_shape = [U.shape[1] for U in factors[m:]] |
|
|
| out1 = factors[0] |
| out2 = factors[m] |
|
|
| for i in range(1, m): |
| out1 = torch.tensordot(out1, factors[i], dims=[[-1], [0]]) |
| out2 = torch.tensordot(out2, factors[i + m], dims=[[-1], [0]]) |
|
|
| out1 = out1.view(np.prod(input_shape), -1) |
| out2 = out2.view(-1, np.prod(output_shape)) |
|
|
| output = torch.tensordot(out1[inds, :], out2, dims=[[-1], [0]]) |
| return output |
|
|
| def ttm_select_ind(self, factors, inds, ind2coord_dict, use_unique=False): |
| input_shape = [U.shape[1] for U in factors] |
| output_shape = [U.shape[2] for U in factors] |
| m = len(self.shape[0]) // 2 |
|
|
| out1 = factors[0] |
| out2 = factors[m] |
|
|
| for i in range(1, m): |
| out1 = torch.tensordot(out1, factors[i], dims=[[-1], [0]]) |
| out2 = torch.tensordot(out2, factors[i + m], dims=[[-1], [0]]) |
|
|
| len1 = len(out1.shape) |
| len2 = len(out2.shape) |
| if len1 > 4: |
| out1 = ( |
| out1.permute( |
| *( |
| [0] |
| + [2 * i + 1 for i in range((len1 - 2) // 2)] |
| + [2 * (i + 1) for i in range((len1 - 2) // 2)] |
| + [len1 - 1] |
| ) |
| ) |
| .contiguous() |
| .view(np.prod(input_shape[:m]), np.prod(output_shape[:m]), -1) |
| ) |
| if len2 > 4: |
| out2 = ( |
| out2.permute( |
| *( |
| [0] |
| + [2 * i + 1 for i in range((len2 - 2) // 2)] |
| + [2 * (i + 1) for i in range((len2 - 2) // 2)] |
| + [len2 - 1] |
| ) |
| ) |
| .contiguous() |
| .view(-1, np.prod(input_shape[m:]), np.prod(output_shape[m:])) |
| ) |
|
|
| |
| if use_unique: |
| inds_cpu = inds |
| inds_unique = torch.unique(inds_cpu) |
| self.dict_tensor[inds_unique] = self.dict_range[0 : inds_unique.shape[0]] |
|
|
| targets = ind2coord_dict[inds_unique, :] |
|
|
| out1 = out1[targets[:, 0], :, :] |
| out2 = out2[:, targets[:, 1], :] |
|
|
| out = torch.einsum("abc,cad->abd", out1, out2).flatten(start_dim=1) |
| out = out[self.dict_tensor[inds_cpu], :] |
| else: |
| |
| |
| targets = ind2coord_dict[inds, :] |
|
|
| out1 = out1[targets[:, 0], :, :] |
| out2 = out2[:, targets[:, 1], :] |
| out = torch.einsum("abc,cad->abd", out1, out2).flatten(start_dim=1) |
|
|
| return out |
|
|
| def forward(self, x, config_forward=None): |
| xshape = list(x.shape) |
| xshape_new = xshape + [ |
| self.out_features, |
| ] |
| x = torch.flatten(x) |
|
|
| if config_forward == None: |
| factors = self.tensor.get_factors(prune_mask=False) |
| else: |
| factors = self.tensor.get_factors( |
| prune_mask=config_forward.prune_mask, threshold=config_forward.threshold |
| ) |
|
|
| if self.factorization == "ttm": |
| rows = self.ttm_select_ind(factors, x, self.ind2coord, use_unique=False) |
| elif self.factorization == "tt": |
| rows = self.tt_select_ind(factors, x) |
| rows = rows.view(*xshape_new) |
|
|
| return rows |
|
|
|
|
| class TensorizedLinear(nn.Module): |
| def __init__(self, in_features, out_features, config, bias=True): |
| """ |
| config has following attributes: |
| shape: the shape of the tensor |
| ranks: either a number or a list of numbers to specify the ranks |
| set_scale_factors: True or False |
| """ |
|
|
| super(TensorizedLinear, self).__init__() |
|
|
| self.in_features = in_features |
| self.out_features = out_features |
| self.shape = config.shape |
| self.factorization = config.factorization |
| self.config_forward = getattr(config, "config_forward", None) |
|
|
| if getattr(config, "lr_act", None) is not None: |
| self.lr_act = ACT2FN[config.lr_act]() |
|
|
| target_stddev = np.sqrt(1 / (self.in_features + self.out_features)) |
| config_tensor = config_class( |
| shape=config.shape, |
| ranks=config.ranks, |
| target_sdv=target_stddev, |
| build_rank_parameters=config.build_rank_parameters, |
| per_decomp_rank_ratio_limit=config.per_decomp_rank_ratio_limit, |
| batch_size=config.batch_size if config.factorization == "btt" else None, |
| ) |
|
|
| |
| if config.factorization == "tt": |
| self.tensor = TensorTrain(config_tensor) |
| if getattr(config, "tied_weight", False): |
| self.fwd_func = ( |
| self.forward_tt_full_precision_transpose_with_original_factors |
| ) |
| elif getattr(config, "amp", False): |
| self.fwd_func = self.forward_tt_with_amp |
| else: |
| self.fwd_func = self.forward_tt_full_precision |
| elif config.factorization == "ttm": |
| self.tensor = TensorTrainMatrix(config_tensor) |
| if getattr(config, "tied_weight", False): |
| self.fwd_func = ( |
| self.forward_ttm_full_precision_transpose_with_original_factors |
| ) |
| else: |
| self.fwd_func = self.forward_ttm_full_precision |
| elif config.factorization == "btt": |
| self.tensor = BatchTensorTrain(config_tensor) |
| self.fwd_func = self.forward_btt_full_precision |
| elif config.factorization == "gcp_tt": |
| assert hasattr(config, "gcp_idx") |
| self.tensor = GCPTensorTrain(config_tensor, config.gcp_idx) |
| self.fwd_func = ( |
| self.forward_gcp_tt_left |
| if not config.gcp_idx |
| else self.forward_gcp_tt_right |
| ) |
| else: |
| raise ValueError( |
| f"Factorization type {config.factorization} not supported." |
| ) |
|
|
| if bias == False: |
| self.register_parameter("bias", None) |
| else: |
| stdv = 1.0 / math.sqrt(out_features) |
| |
| |
| self.bias = torch.nn.Parameter(torch.randn(out_features)) |
| self.bias.data.uniform_(-stdv, stdv) |
|
|
| if hasattr(config, "set_scale_factors") and config.set_scale_factors: |
| self.set_scale_factors() |
|
|
| def set_scale_factors( |
| self, |
| scale_w=1.0, |
| scale_input=1.0, |
| scale_intermediate=1.0, |
| scale_dy=1.0, |
| scale_x=1.0, |
| scale_out=1.0, |
| ): |
| self.scales = torch.nn.ParameterList() |
| self.scale_factors = torch.nn.ParameterList() |
|
|
| if not isinstance(scale_w, list): |
| scale_w = [scale_w] * self.tensor.order |
| for s in scale_w: |
| self.scale_factors.append(torch.nn.Parameter(torch.tensor(s))) |
|
|
| self.scale_input = torch.nn.Parameter(torch.tensor(scale_input)) |
| self.scale_intermediate = torch.nn.Parameter(torch.tensor(scale_intermediate)) |
| self.scale_dy = torch.nn.Parameter(torch.tensor(scale_dy)) |
| self.scale_x = torch.nn.Parameter(torch.tensor(scale_x)) |
| self.scale_out = torch.nn.Parameter(torch.tensor(scale_out)) |
|
|
| self.scales.append(self.scale_input) |
| self.scales.append(self.scale_intermediate) |
| self.scales.append(self.scale_dy) |
| self.scales.append(self.scale_x) |
| self.scales.append(self.scale_out) |
|
|
| def forward(self, input): |
| """ |
| config_forward: |
| prune_mask: True or False. Use prune mask or not |
| threshold: float number. The threshold to clip rank_parameters to 0 |
| quantized: 0: full precision. 1: quantization-aware training. 2: low-precision training. |
| if quantized: |
| rep: INT or FLOAT. quantization type |
| bit_input/factors/intermediate/out: bits for each part |
| rounding: stochastic or nearest. Rounding type |
| """ |
| if self.config_forward is None: |
| factors = self.tensor.get_factors(prune_mask=False) |
| else: |
| factors = self.tensor.get_factors( |
| prune_mask=self.config_forward.prune_mask, |
| threshold=self.config_forward.threshold, |
| ) |
|
|
| out = self.fwd_func(input, factors) |
| if self.bias is not None: |
| out += self.bias |
|
|
| return out |
|
|
| def forward_btt_full_precision(self, input_mat, factors): |
| m = len(factors) // 2 |
| N = len(input_mat.shape) |
|
|
| out = factors[0] |
| r1 = out.shape[-1] |
| bz = out.shape[0] |
|
|
| output = factors[m] |
| r3 = output.shape[-1] |
|
|
| for i in range(1, m): |
| U = factors[i] |
| V = factors[i + m] |
| r2 = U.shape[-1] |
| r4 = V.shape[-1] |
|
|
| out = torch.bmm(out.view(bz, -1, r1), U.view(bz, r1, -1)) |
| output = torch.bmm(output.view(bz, -1, r3), V.view(bz, r3, -1)) |
|
|
| r1 = r2 |
| r3 = r4 |
|
|
| out = torch.einsum("abc, dce -> adbe", input_mat, out.view(bz, -1, r1)) |
| output = torch.einsum("abcd, bde -> abce", out, output.view(bz, r1, -1)) |
|
|
| return output |
|
|
| def forward_gcp_tt_left(self, input_mat: torch.Tensor, factors: List[torch.Tensor]): |
| m = len(factors) |
| N = len(input_mat.shape) |
|
|
| input_mat = input_mat.view(list(input_mat.shape[0 : N - 1]) + self.shape[:m]) |
|
|
| out = factors[0].squeeze() |
|
|
| for i in range(1, m): |
| U = factors[i] |
| out = torch.tensordot(out, U, [[-1], [0]]) |
|
|
| out = torch.tensordot(input_mat, out, [list(range(-m, 0)), list(range(0, m))]) |
| assert len(out.shape) == N |
|
|
| return out |
|
|
| def forward_gcp_tt_right( |
| self, input_mat: torch.Tensor, factors: List[torch.Tensor] |
| ): |
| m = len(factors) |
| N = len(input_mat.shape) |
|
|
| out = factors[0] |
|
|
| for i in range(1, m): |
| U = factors[i] |
| out = torch.tensordot(out, U, [[-1], [0]]) |
|
|
| out = torch.tensordot(input_mat, out, [[-1], [0]]).flatten(start_dim=N - 1) |
| assert len(out.shape) == N |
|
|
| return out |
|
|
| def forward_tt_full_precision(self, input_mat, factors): |
| m = len(factors) // 2 |
| N = len(input_mat.shape) |
|
|
| input_mat = torch.reshape( |
| input_mat, [1] + list(input_mat.shape[0 : N - 1]) + self.shape[:m] |
| ) |
|
|
| out = factors[0] |
|
|
| out = torch.squeeze(out) |
| output = factors[m] |
|
|
| for i in range(1, m): |
| U = factors[i] |
| V = factors[i + m] |
|
|
| out = torch.tensordot(out, U, [[-1], [0]]) |
| output = torch.tensordot(output, V, [[-1], [0]]) |
|
|
| out = torch.tensordot( |
| input_mat, out, [list(range(N, N + m)), list(range(0, m))] |
| ) |
|
|
| N = len(out.shape) |
|
|
| if hasattr(self, "lr_act"): |
| out = self.lr_act(out) |
|
|
| output = torch.tensordot(out, output, [[-1], [0]]) |
|
|
| output = ( |
| torch.flatten(output, start_dim=N - 1, end_dim=-1).squeeze_(0).squeeze_(-1) |
| ) |
|
|
| return output |
|
|
| def forward_tt_with_amp(self, input_mat, factors): |
| m = len(factors) // 2 |
| N = len(input_mat.shape) |
|
|
| out = factors[0] |
|
|
| out = torch.squeeze(out) |
| output = factors[m] |
|
|
| for i in range(1, m): |
| U = factors[i] |
| V = factors[i + m] |
| r_0 = U.shape[0] |
| r_1 = V.shape[0] |
|
|
| out = torch.matmul(out.view(-1, r_0), U.view(r_0, -1)) |
| output = torch.matmul(output.view(-1, r_1), V.view(r_1, -1)) |
|
|
| out = torch.matmul(input_mat, out.view(input_mat.shape[-1], -1)) |
|
|
| if hasattr(self, "lr_act"): |
| out = self.lr_act(out) |
| |
| output = ( |
| torch.matmul(out, output.view(out.shape[-1], -1)).squeeze_(0).squeeze_(-1) |
| ) |
|
|
| return output |
|
|
| def forward_tt_full_precision_transpose_with_original_factors( |
| self, input_mat, factors |
| ): |
| m = len(factors) // 2 |
| N = len(input_mat.shape) |
|
|
| input_mat = torch.reshape( |
| input_mat, [1] + list(input_mat.shape[0 : N - 1]) + self.shape[m:] |
| ) |
|
|
| out = factors[m] |
| output = factors[0] |
|
|
| for i in range(1, m): |
| U = factors[i + m] |
| V = factors[i] |
|
|
| out = torch.tensordot(out, U, [[-1], [0]]) |
| output = torch.tensordot(output, V, [[-1], [0]]) |
|
|
| out = torch.tensordot( |
| input_mat, |
| out.squeeze(), |
| dims=[list(range(N, N + m)), list(range(1, 1 + m))], |
| ) |
|
|
| N = len(out.shape) |
| output = torch.tensordot(out, output.squeeze(), dims=[[-1], [-1]]) |
|
|
| output = torch.flatten(output, start_dim=N - 1, end_dim=-1) |
| output = torch.squeeze(output) |
|
|
| return output |
|
|
| def forward_ttm_full_precision(self, input_mat, factors): |
| M = len(factors) |
| N = len(input_mat.shape) |
|
|
| input_mat = torch.reshape( |
| input_mat, list(input_mat.shape[0 : N - 1]) + self.shape[0] + [1] |
| ) |
| output = input_mat |
|
|
| for i in range(M): |
| output = torch.tensordot(input_mat, factors[i], [[N - 1, -1], [1, 0]]) |
|
|
| output = torch.flatten(output, start_dim=N - 1, end_dim=-1) |
|
|
| return output |
|
|
| |
| def forward_ttm_full_precision_transpose_with_transposed_factors( |
| self, input_mat, factors |
| ): |
| M = len(factors) |
| N = len(input_mat.shape) |
|
|
| input_mat = torch.reshape( |
| input_mat, list(input_mat.shape[0 : N - 1]) + self.shape[0] + [1] |
| ) |
| output = input_mat |
|
|
| for i in range(M): |
| output = torch.tensordot(output, factors[i], [[N + M - 2, -1], [1, 0]]) |
| M -= 1 |
|
|
| output = output.permute( |
| *([0] + [i for i in range(N - 1, N - 1 + len(factors))][::-1] + [-1]) |
| ) |
| output = torch.flatten(output, start_dim=N - 1, end_dim=-1) |
|
|
| return output |
|
|
| def forward_ttm_full_precision_transpose_with_original_factors( |
| self, input_mat, factors |
| ): |
| M = len(factors) |
| N = len(input_mat.shape) |
|
|
| input_mat = torch.reshape( |
| input_mat, list(input_mat.shape[0 : N - 1]) + self.shape[1] + [1] |
| ) |
| output = input_mat |
|
|
| for i in range(M): |
| |
| output = torch.tensordot( |
| output, factors[::-1][i], [[N + M - 2, -2 if i else -1], [2, -1]] |
| ) |
| M -= 1 |
|
|
| output = output.permute( |
| *( |
| list(range(N - 1)) |
| + [-1] |
| + list(range(N - 1, N - 2 + len(factors)))[::-1] |
| + [-2] |
| ) |
| ) |
| output = torch.flatten(output, start_dim=N - 1, end_dim=-1) |
|
|
| return output |
|
|