File size: 3,361 Bytes
4ca4e4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
import os
from models import (
        LSTM,
        GPTNeoXAlibiForCausalLM,
        GPTNeoXHardAlibiForCausalLM,
        GPTNeoXNoPEForCausalLM,
        GPTNeoXForCausalLM,
        MambaForCausalLM,
        HybridNoPEForCausalLM,
        HybridForCausalLM,
        HybridNoPEConfig,
        HybridConfig
    )

from transformers import GPTNeoXConfig
from transformers import MambaConfig
from transformers import AutoModelForCausalLM

def get_model(args, tokenizer):
    if args.model in ["T_nope","T_rope","T_alibi"]:
        config = GPTNeoXConfig(
            bos_token_id=0,
            eos_token_id=0,
            hidden_size=args.hidden_size,
            intermediate_size=args.hidden_size*4,
            num_attention_heads=args.heads,
            num_hidden_layers=args.layers,
            vocab_size=len(tokenizer),
        )
    elif args.model == "T_hard_alibi":
        config = GPTNeoXConfig(
            bos_token_id=0,
            eos_token_id=0,
            hidden_size=args.hidden_size,
            intermediate_size=args.hidden_size*4,
            num_attention_heads=args.heads,
            num_hidden_layers=args.layers,
            num_masked_heads=args.num_masked_heads,
            vocab_size=len(tokenizer),
        )
    elif args.model == "mamba":
        config = MambaConfig(
            hidden_size=args.hidden_size,
            d_model=args.hidden_size,
            n_layer=args.layers,
            ssm_cfg={"d_state": args.state_dim},
            vocab_size=len(tokenizer),
        )
    elif args.model == "hybrid":
        config = HybridConfig(
            bos_token_id=0,
            eos_token_id=0,
            hidden_size=args.hidden_size,
            intermediate_size=args.hidden_size*4,
            num_attention_heads=args.heads,
            num_hidden_layers=args.layers,
            d_model=args.hidden_size,
            n_layer=args.layers,
            ssm_cfg={"d_state": args.state_dim},
            vocab_size=len(tokenizer),
        )
    elif args.model == "hybrid_nope":
        config = HybridNoPEConfig(
            bos_token_id=0,
            eos_token_id=0,
            hidden_size=args.hidden_size,
            intermediate_size=args.hidden_size*4,
            num_attention_heads=args.heads,
            num_hidden_layers=args.layers,
            d_model=args.hidden_size,
            n_layer=args.layers,
            ssm_cfg={"d_state": args.state_dim},
            vocab_size=len(tokenizer),
        )


    
    if args.model=="T_rope":
        model = GPTNeoXForCausalLM(config)
    elif args.model=="T_nope":
        model = GPTNeoXNoPEForCausalLM(config)
    elif args.model=="T_alibi":
        model = GPTNeoXAlibiForCausalLM(config)
    elif args.model=="T_hard_alibi":
        model = GPTNeoXHardAlibiForCausalLM(config)
    elif args.model=="mamba":
        model = MambaForCausalLM(config)
    elif args.model=="lstm":
        model = LSTM(
            embedding_dim=args.hidden_size,
            vocab_size=len(tokenizer),
            num_layers=args.layers,
            dropout_rate=0.65
        )
    elif args.model=="hybrid":
        model = HybridForCausalLM(config)
    elif args.model=="hybrid_nope":
        model = HybridNoPEForCausalLM(config)

    if args.model == "pretrained":
        model = AutoModelForCausalLM.from_pretrained(args.pretrain_model)
        
    return model