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

Class QuantLlamaAttentionFused

inference/modules/fused_attn.py:166–301  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

164
165
166class QuantLlamaAttentionFused(nn.Module):
167 def __init__(self, hidden_size, num_heads, qkv_layer, o_proj, dev, args):
168 super().__init__()
169
170 self.args = args
171 self.n_local_heads = args.num_attention_heads
172 self.hidden_size = args.hidden_size
173 self.num_heads = args.num_attention_heads
174 self.head_dim = self.hidden_size // self.num_heads
175
176 self.num_key_value_heads = args.num_key_value_heads
177 self.num_key_value_groups = self.num_heads // self.num_key_value_heads
178 self.max_position_embeddings = args.max_position_embeddings
179 # self.rope_theta = args.rope_theta
180
181 self.qkv_proj = qkv_layer
182 self.o_proj = o_proj
183
184 kv_max_seq_len = min(max_seq_len, args.max_position_embeddings)
185
186 # following fastertransformer definition
187 self.cache_v = (
188 torch.zeros(
189 (
190 max_batch_size,
191 self.num_key_value_heads,
192 # args.max_position_embeddings,
193 kv_max_seq_len,
194 self.head_dim,
195 )
196 )
197 .to(dev)
198 .half()
199 ) # added to half
200 # 8: pack 8 fp16 in FT, if fp32 then use 4
201 self.cache_k = (
202 torch.zeros(
203 (
204 max_batch_size,
205 self.num_key_value_heads,
206 self.head_dim // 8,
207 # args.max_position_embeddings,
208 kv_max_seq_len,
209 8,
210 )
211 )
212 .to(dev)
213 .half()
214 ) # added to half
215
216 # dummy
217 self.rotary_emb = QuantLlamaRotaryEmbedding(
218 self.head_dim, max_position_embeddings=2048, device="cuda:0"
219 )
220
221 def forward(
222 self,
223 x: torch.Tensor,

Callers 1

make_quant_attnFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected