kgrabko commited on
Commit
4d33ecd
·
verified ·
1 Parent(s): 036d185

Create chat_405b_deepspeed.py

Browse files
Files changed (1) hide show
  1. chat_405b_deepspeed.py +90 -0
chat_405b_deepspeed.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #%%writefile chat_405b.py
2
+ # ==============================================================================
3
+ # COPYRIGHT (C) 2025-2026 KONSTANTIN VLADIMIROVICH GRABKO. ALL RIGHTS RESERVED.
4
+ # JiRack 405B Chat with DeepSpeed ZeRO-3 Offload
5
+ # ==============================================================================
6
+
7
+ import torch
8
+ import os
9
+ import deepspeed
10
+ from transformers import AutoTokenizer
11
+ from JiRackTernaryPyTorch_405b import JiRackTernaryModel
12
+
13
+ # --- ПУТИ ---
14
+ LOCAL_PATH = "/content/drive/MyDrive/JiRack_405B_SFT_Steps_2/step_4500"
15
+ TOKENIZER_ID = "meta-llama/Meta-Llama-3.1-405B-Instruct"
16
+
17
+ # --- DEEPSPEED CONFIG ---
18
+ ds_config = {
19
+ "fp16": {"enabled": True},
20
+ "zero_optimization": {
21
+ "stage": 3,
22
+ "offload_param": {"device": "cpu", "pin_memory": True},
23
+ "overlap_comm": True,
24
+ "contiguous_gradients": True,
25
+ "stage3_max_live_parameters": 1e8,
26
+ "stage3_max_reuse_distance": 1e8,
27
+ "stage3_prefetch_bucket_size": 5e7,
28
+ "stage3_param_persistence_threshold": 1e5,
29
+ },
30
+ "steps_per_print": 2000,
31
+ "train_batch_size": 1,
32
+ }
33
+
34
+ def main():
35
+ print("[*] Loading tokenizer...")
36
+ tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_ID)
37
+ if tokenizer.pad_token is None:
38
+ tokenizer.pad_token = tokenizer.eos_token
39
+
40
+ print("[*] Initializing JiRack 405B with DeepSpeed ZeRO-3...")
41
+
42
+ # Инициализация модели внутри контекста DeepSpeed Zero
43
+ # Это позволяет загружать веса по частям, не забивая RAM мгновенно
44
+ with deepspeed.zero.Init(config_dict_or_path=ds_config):
45
+ model = JiRackTernaryModel.from_pretrained(
46
+ LOCAL_PATH,
47
+ torch_dtype=torch.float16,
48
+ trust_remote_code=True
49
+ )
50
+
51
+ # Инициализация движка DeepSpeed
52
+ ds_engine = deepspeed.initialize(model=model, config_params=ds_config)[0]
53
+ ds_engine.eval()
54
+
55
+ print("\n✅ JiRack 405B is ready (ZeRO-3 Offload Active).")
56
+ print("Type 'exit' to quit.\n")
57
+
58
+ while True:
59
+ try:
60
+ user_input = input("User: ")
61
+ if user_input.lower() in ["exit", "quit"]: break
62
+ if not user_input.strip(): continue
63
+
64
+ print("JiRack thinking...", end="\r")
65
+
66
+ # Форматируем промпт
67
+ prompt = f"<|begin_of_text|><|start_header_id|>user<|end_header_id|>\n\n{user_input}<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n"
68
+ inputs = tokenizer(prompt, return_tensors="pt").to(torch.cuda.current_device())
69
+
70
+ with torch.no_grad():
71
+ # Используем .module для доступа к оригинальному методу generate
72
+ outputs = ds_engine.module.generate(
73
+ **inputs,
74
+ max_new_tokens=128,
75
+ do_sample=True,
76
+ temperature=0.7,
77
+ top_p=0.9,
78
+ eos_token_id=[tokenizer.eos_token_id, 128001, 128008, 128009]
79
+ )
80
+
81
+ response = tokenizer.decode(outputs[0][inputs['input_ids'].shape[-1]:], skip_special_tokens=True)
82
+ print(f"JiRack: {response}\n")
83
+
84
+ except KeyboardInterrupt:
85
+ break
86
+ except Exception as e:
87
+ print(f"\n[!] Error: {e}")
88
+
89
+ if __name__ == "__main__":
90
+ main()