MCPcopy Create free account
hub / github.com/pytorch/examples / __init__

Method __init__

distributed/tensor_parallelism/llama2_model.py:165–183  ·  view source on GitHub ↗
(self, model_args: ModelArgs)

Source from the content-addressed store, hash-verified

163 """
164
165 def __init__(self, model_args: ModelArgs):
166 super().__init__()
167 self.n_heads = model_args.n_heads
168 self.n_kv_heads = (
169 model_args.n_heads
170 if model_args.n_kv_heads is None
171 else model_args.n_kv_heads
172 )
173 self.n_rep = self.n_heads // self.n_kv_heads
174 self.head_dim = model_args.dim // model_args.n_heads
175
176 self.wq = nn.Linear(
177 model_args.dim, model_args.n_heads * self.head_dim, bias=False
178 )
179 self.wk = nn.Linear(model_args.dim, self.n_kv_heads * self.head_dim, bias=False)
180 self.wv = nn.Linear(model_args.dim, self.n_kv_heads * self.head_dim, bias=False)
181 self.wo = nn.Linear(
182 model_args.n_heads * self.head_dim, model_args.dim, bias=False
183 )
184
185 def init_weights(self, init_std: float):
186 for linear in (self.wq, self.wk, self.wv):

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected