mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
einops for tensor manipulation
This commit is contained in:
parent
9649a11bc2
commit
4537e1c11b
1 changed files with 21 additions and 7 deletions
|
|
@ -3,6 +3,7 @@ from dataclasses import dataclass, field
|
|||
from enum import Enum
|
||||
from functools import partial
|
||||
from typing import Any, Optional, Tuple, Union
|
||||
from einops import rearrange, repeat
|
||||
|
||||
import torch
|
||||
from jaxtyping import Float, Integer
|
||||
|
|
@ -107,17 +108,28 @@ class Pooler(nn.Module):
|
|||
|
||||
# [bs, feature_dim]
|
||||
x = self.ln(self.feature_proj(self.pool_fn(features, attn_mask).float()))
|
||||
# [bs, feature_dim] -> [bs, num_modules, num_layers, feature_dim]
|
||||
x = x.expand(self.num_modules, self.num_layers, -1, -1)
|
||||
x = x.permute((2, 0, 1, 3))
|
||||
x = repeat(
|
||||
x,
|
||||
"bs d -> bs n_modules n_layers d",
|
||||
n_modules=self.num_modules,
|
||||
n_layers=self.num_layers,
|
||||
)
|
||||
|
||||
layer_embs = self.layer_embs(self.layer_indices) # [num_layers, d]
|
||||
# [num_layers, d] -> [bs, num_modules, num_layers, d]
|
||||
layer_embs = layer_embs.expand(bs, self.num_modules, -1, -1)
|
||||
layer_embs = repeat(
|
||||
layer_embs,
|
||||
"n_layers d -> bs n_modules n_layers d",
|
||||
bs=bs,
|
||||
n_modules=self.num_modules,
|
||||
)
|
||||
|
||||
module_embs = self.module_embs(self.module_indices) # [num_modules, d]
|
||||
# [num_modules, d] -> [bs, num_modules, num_layers, d]
|
||||
module_embs = module_embs.expand(bs, self.num_layers, -1, -1).transpose(1, 2)
|
||||
module_embs = repeat(
|
||||
module_embs,
|
||||
"n_modules d -> bs n_modules n_layers d",
|
||||
bs=bs,
|
||||
n_layers=self.num_layers,
|
||||
)
|
||||
|
||||
emb = torch.cat([x, layer_embs, module_embs], dim=3)
|
||||
return self.mlp(self.mixer(emb))
|
||||
|
|
@ -234,6 +246,8 @@ class ModulatedPretrainedModel(nn.Module):
|
|||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# set torch randomness seed
|
||||
torch.manual_seed(42)
|
||||
model_name = "meta-llama/Llama-3.2-1B-Instruct"
|
||||
base_model, tokenizer = get_model_and_tokenizer(
|
||||
model_name,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue