MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / QuantLlamaAttention

Class QuantLlamaAttention

inference/modules/fused_attn.py:79–163  ·  view source on GitHub ↗

Multi-headed attention from 'Attention Is All You Need' paper

Source from the content-addressed store, hash-verified

77
78
79class QuantLlamaAttention(nn.Module):
80 """Multi-headed attention from 'Attention Is All You Need' paper"""
81
82 def __init__(self, hidden_size, num_heads, qkv_proj, o_proj, dev):
83 super().__init__()
84 self.hidden_size = hidden_size
85 self.num_heads = num_heads
86 self.head_dim = hidden_size // num_heads
87
88 if (self.head_dim * num_heads) != self.hidden_size:
89 raise ValueError(
90 f"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}"
91 f" and `num_heads`: {num_heads})."
92 )
93 self.qkv_proj = qkv_proj
94 self.o_proj = o_proj
95 self.rotary_emb = QuantLlamaRotaryEmbedding(
96 self.head_dim, max_position_embeddings=2048, device=dev
97 )
98
99 def forward(
100 self,
101 hidden_states,
102 past_key_value=None,
103 attention_mask=None,
104 position_ids=None,
105 output_attentions=False,
106 use_cache=False,
107 ):
108 """Input shape: Batch x Time x Channel"""
109
110 bsz, q_len, _ = hidden_states.size()
111
112 qkv_states = self.qkv_proj(hidden_states)
113 qkv_states = qkv_states.view(bsz, q_len, 3, self.num_heads, self.head_dim)
114
115 # This updates the query and key states in-place, saving VRAM.
116 query_states, key_states, value_states = torch.split(qkv_states, 1, dim=2)
117 query_states, key_states = self.rotary_emb(
118 query_states, key_states, position_ids
119 )
120
121 del qkv_states
122 query_states = query_states.view(
123 bsz, q_len, self.num_heads, self.head_dim
124 ).transpose(1, 2)
125 key_states = key_states.view(
126 bsz, q_len, self.num_heads, self.head_dim
127 ).transpose(1, 2)
128 value_states = value_states.view(
129 bsz, q_len, self.num_heads, self.head_dim
130 ).transpose(1, 2)
131
132 is_causal = past_key_value is None
133
134 kv_seq_len = q_len
135 if past_key_value is not None:
136 kv_seq_len += past_key_value[0].shape[-2]

Callers 1

make_quant_attnFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected