mirror of
https://github.com/jingyaogong/minimind.git
synced 2026-10-06 09:07:31 +00:00
update config
This commit is contained in:
@@ -51,9 +51,9 @@ def push_to_hf(export_model_path):
|
||||
if __name__ == '__main__':
|
||||
lm_config = LMConfig(dim=512, n_layers=8, max_seq_len=8192, use_moe=False)
|
||||
|
||||
torch_path = f"../out/reason_{lm_config.dim}{'_moe' if lm_config.use_moe else ''}.pth"
|
||||
torch_path = f"../out/rlhf_{lm_config.dim}{'_moe' if lm_config.use_moe else ''}.pth"
|
||||
|
||||
transformers_path = '../MiniMind2-Small-R1'
|
||||
transformers_path = '../MiniMind2-Small'
|
||||
|
||||
# convert torch to transformers model
|
||||
convert_torch2transformers(torch_path, transformers_path)
|
||||
|
||||
@@ -26,7 +26,7 @@ def init_model(args):
|
||||
tokenizer = AutoTokenizer.from_pretrained('../model/minimind_tokenizer')
|
||||
if args.load == 0:
|
||||
moe_path = '_moe' if args.use_moe else ''
|
||||
modes = {0: 'pretrain', 1: 'full_sft', 2: 'full_dist', 3: 'rlhf'}
|
||||
modes = {0: 'pretrain', 1: 'full_sft', 2: 'rlhf', 3: 'reason'}
|
||||
ckp = f'../{args.out_dir}/{modes[args.model_mode]}_{args.dim}{moe_path}.pth'
|
||||
|
||||
model = MiniMindLM(LMConfig(
|
||||
@@ -156,7 +156,7 @@ if __name__ == "__main__":
|
||||
parser.add_argument('--max_seq_len', default=8192, type=int)
|
||||
parser.add_argument('--use_moe', default=False, type=bool)
|
||||
parser.add_argument('--load', default=0, type=int, help="0: 从原生torch权重,1: 利用transformers加载")
|
||||
parser.add_argument('--model_mode', default=1, type=int, help="0: 预训练模型,1: SFT-Chat模型,2: RLHF-Chat模型")
|
||||
parser.add_argument('--model_mode', default=1, type=int, help="0: 预训练模型,1: SFT-Chat模型,2: RLHF-Chat模型,3: Reason模型")
|
||||
|
||||
device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
model, tokenizer = init_model(parser.parse_args())
|
||||
|
||||
Reference in New Issue
Block a user