MCPcopy Create free account
hub / github.com/kyegomez/BitNet / __init__

Method __init__

bitnet/bit_llama.py:178–251  ·  view source on GitHub ↗

Initialize the Attention module. Args: args (ModelArgs): Model configuration parameters. Attributes: n_kv_heads (int): Number of key and value heads. n_local_heads (int): Number of local query heads. n_local_kv_heads (int): N

(self, args: ModelArgs)

Source from the content-addressed store, hash-verified

176 """Multi-head attention module."""
177
178 def __init__(self, args: ModelArgs):
179 """
180 Initialize the Attention module.
181
182 Args:
183 args (ModelArgs): Model configuration parameters.
184
185 Attributes:
186 n_kv_heads (int): Number of key and value heads.
187 n_local_heads (int): Number of local query heads.
188 n_local_kv_heads (int): Number of local key and value heads.
189 n_rep (int): Number of repetitions for local heads.
190 head_dim (int): Dimension size of each attention head.
191 wq (ColumnParallelLinear): Linear transformation for queries.
192 wk (ColumnParallelLinear): Linear transformation for keys.
193 wv (ColumnParallelLinear): Linear transformation for values.
194 wo (RowParallelLinear): Linear transformation for output.
195 cache_k (torch.Tensor): Cached keys for attention.
196 cache_v (torch.Tensor): Cached values for attention.
197
198 """
199 super().__init__()
200 self.n_kv_heads = args.n_heads if args.n_kv_heads is None else args.n_kv_heads
201 model_parallel_size = fs_init.get_model_parallel_world_size()
202 self.n_local_heads = args.n_heads // model_parallel_size
203 self.n_local_kv_heads = self.n_kv_heads // model_parallel_size
204 self.n_rep = self.n_local_heads // self.n_local_kv_heads
205 self.head_dim = args.dim // args.n_heads
206
207 self.wq = ColumnParallelLinear(
208 args.dim,
209 args.n_heads * self.head_dim,
210 bias=False,
211 gather_output=False,
212 init_method=lambda x: x,
213 )
214 self.wk = ColumnParallelLinear(
215 args.dim,
216 self.n_kv_heads * self.head_dim,
217 bias=False,
218 gather_output=False,
219 init_method=lambda x: x,
220 )
221 self.wv = ColumnParallelLinear(
222 args.dim,
223 self.n_kv_heads * self.head_dim,
224 bias=False,
225 gather_output=False,
226 init_method=lambda x: x,
227 )
228 self.wo = RowParallelLinear(
229 args.n_heads * self.head_dim,
230 args.dim,
231 bias=False,
232 input_is_parallel=True,
233 init_method=lambda x: x,
234 )
235

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected