| 1947 | """ |
| 1948 | |
| 1949 | def __init__( |
| 1950 | self, |
| 1951 | hidden_size: int, |
| 1952 | cross_attention_dim: Optional[int] = None, |
| 1953 | rank: int = 4, |
| 1954 | network_alpha: Optional[int] = None, |
| 1955 | ): |
| 1956 | super().__init__() |
| 1957 | |
| 1958 | self.hidden_size = hidden_size |
| 1959 | self.cross_attention_dim = cross_attention_dim |
| 1960 | self.rank = rank |
| 1961 | |
| 1962 | self.to_q_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha) |
| 1963 | self.add_k_proj_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha) |
| 1964 | self.add_v_proj_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha) |
| 1965 | self.to_k_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha) |
| 1966 | self.to_v_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha) |
| 1967 | self.to_out_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha) |
| 1968 | |
| 1969 | def __call__(self, attn: Attention, hidden_states: torch.FloatTensor, *args, **kwargs) -> torch.FloatTensor: |
| 1970 | self_cls_name = self.__class__.__name__ |