gcharanteja commited on
Commit
4117f93
Β·
1 Parent(s): 86ada7e
Files changed (1) hide show
  1. fine_tune_llama32_1b.py +106 -20
fine_tune_llama32_1b.py CHANGED
@@ -1,13 +1,99 @@
1
  import torch
 
 
 
 
 
2
  from datasets import load_dataset
3
  from transformers import (
4
  AutoModelForCausalLM,
5
  AutoTokenizer,
6
  BitsAndBytesConfig,
7
  )
8
- from peft import LoraConfig, prepare_model_for_kbit_training, get_peft_model
9
  from trl import SFTTrainer, SFTConfig
10
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
11
  # ========================= CONFIG =========================
12
  model_name = "meta-llama/Llama-3.2-1B-Instruct"
13
  dataset_name = "mlabonne/FineTome-100k"
@@ -22,6 +108,9 @@ max_steps = None # set a number (e.g. 2000) if you want to train o
22
 
23
  output_dir = "./llama32-1b-finetuned-finetome"
24
 
 
 
 
25
  # ====================== LOAD MODEL (4-bit QLoRA) ======================
26
  bnb_config = BitsAndBytesConfig(
27
  load_in_4bit=True,
@@ -35,9 +124,14 @@ model = AutoModelForCausalLM.from_pretrained(
35
  quantization_config=bnb_config,
36
  device_map="auto", # automatically puts layers on GPU
37
  trust_remote_code=True,
 
38
  )
39
 
40
- tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
 
 
 
 
41
  tokenizer.pad_token = tokenizer.eos_token
42
 
43
  # Prepare for QLoRA
@@ -57,7 +151,11 @@ model = get_peft_model(model, lora_config)
57
  model.print_trainable_parameters() # should show ~0.5-1% trainable params
58
 
59
  # ====================== LOAD & PREPARE DATASET ======================
60
- dataset = load_dataset(dataset_name, split="train")
 
 
 
 
61
 
62
  # Convert ShareGPT "conversations" β†’ standard messages format
63
  def map_to_messages(example):
@@ -123,25 +221,13 @@ tokenizer.save_pretrained(output_dir)
123
  print(f"βœ… Training finished! LoRA adapter saved to {output_dir}")
124
 
125
 
126
- #model merger from peft import AutoPeftModelForCausalLM
127
-
128
- model = AutoPeftModelForCausalLM.from_pretrained(
129
- output_dir,
130
- device_map="auto",
131
- torch_dtype=torch.bfloat16
132
- )
133
- model = model.merge_and_unload()
134
- model.save_pretrained("llama32-1b-finetuned-merged")
135
- tokenizer.save_pretrained("llama32-1b-finetuned-merged")
136
-
137
-
138
- from peft import AutoPeftModelForCausalLM
139
-
140
  model = AutoPeftModelForCausalLM.from_pretrained(
141
  output_dir,
142
  device_map="auto",
143
- torch_dtype=torch.bfloat16
144
  )
145
  model = model.merge_and_unload()
146
- model.save_pretrained("llama32-1b-finetuned-merged")
147
- tokenizer.save_pretrained("llama32-1b-finetuned-merged")
 
 
1
  import torch
2
+ import os
3
+ import json
4
+ import urllib.request
5
+ import urllib.error
6
+ import getpass
7
  from datasets import load_dataset
8
  from transformers import (
9
  AutoModelForCausalLM,
10
  AutoTokenizer,
11
  BitsAndBytesConfig,
12
  )
13
+ from peft import LoraConfig, prepare_model_for_kbit_training, get_peft_model, AutoPeftModelForCausalLM
14
  from trl import SFTTrainer, SFTConfig
15
 
16
+
17
+ def _fetch_secret_json(url: str, api_key: str, timeout_s: int = 30) -> str:
18
+ req = urllib.request.Request(
19
+ url,
20
+ headers={
21
+ "accept": "application/json",
22
+ "X-API-Key": api_key,
23
+ },
24
+ method="GET",
25
+ )
26
+ try:
27
+ with urllib.request.urlopen(req, timeout=timeout_s) as resp:
28
+ raw = resp.read().decode("utf-8")
29
+ except urllib.error.HTTPError as e:
30
+ body = e.read().decode("utf-8", errors="replace") if hasattr(e, "read") else ""
31
+ raise RuntimeError(f"KeyVault HTTP {e.code}: {body}") from e
32
+ except Exception as e:
33
+ raise RuntimeError(f"KeyVault request failed: {e}") from e
34
+
35
+ try:
36
+ payload = json.loads(raw)
37
+ except json.JSONDecodeError as e:
38
+ raise RuntimeError(f"KeyVault did not return JSON: {raw[:200]}") from e
39
+
40
+ if isinstance(payload, str) and payload.strip():
41
+ return payload.strip()
42
+
43
+ if isinstance(payload, dict):
44
+ for key in ("value", "secret", "token", "hftoken", "hf_token"):
45
+ value = payload.get(key)
46
+ if isinstance(value, str) and value.strip():
47
+ return value.strip()
48
+ if len(payload) == 1:
49
+ value = next(iter(payload.values()))
50
+ if isinstance(value, str) and value.strip():
51
+ return value.strip()
52
+
53
+ raise RuntimeError(f"Unexpected KeyVault JSON shape: {type(payload).__name__}")
54
+
55
+
56
+ def get_hf_token() -> str | None:
57
+ """Get a Hugging Face token without persisting it.
58
+
59
+ Order:
60
+ 1) Use `HF_TOKEN` if set (runtime-only, provided by caller)
61
+ 2) Else fetch from KeyVault URL (default: /secrets/hftoken) using an API key
62
+ - `KEYVAULT_API_KEY` env var, or
63
+ - prompt at runtime (getpass)
64
+
65
+ This intentionally avoids calling `huggingface_hub.login()` to prevent writing
66
+ tokens to disk.
67
+ """
68
+
69
+ token = os.environ.get("HF_TOKEN")
70
+ if token and token.strip():
71
+ return token.strip()
72
+
73
+ keyvault_url = os.environ.get(
74
+ "KEYVAULT_HF_TOKEN_URL",
75
+ "https://maxxcarl-keyvault.hf.space/secrets/hftoken",
76
+ )
77
+
78
+ api_key = os.environ.get("KEYVAULT_API_KEY")
79
+ if not api_key:
80
+ try:
81
+ api_key = getpass.getpass("KeyVault X-API-Key (won't echo): ")
82
+ except Exception:
83
+ api_key = None
84
+
85
+ if api_key and api_key.strip():
86
+ return _fetch_secret_json(keyvault_url, api_key.strip())
87
+
88
+ return None
89
+
90
+
91
+ def _from_pretrained_kwargs_with_token(hf_token: str | None) -> dict:
92
+ if not hf_token:
93
+ return {}
94
+ # Newer HF stacks accept `token=...`; some older call sites still used `use_auth_token`.
95
+ return {"token": hf_token}
96
+
97
  # ========================= CONFIG =========================
98
  model_name = "meta-llama/Llama-3.2-1B-Instruct"
99
  dataset_name = "mlabonne/FineTome-100k"
 
108
 
109
  output_dir = "./llama32-1b-finetuned-finetome"
110
 
111
+ # Fetch an HF token at runtime (no persistence). Required for gated models.
112
+ hf_token = get_hf_token()
113
+
114
  # ====================== LOAD MODEL (4-bit QLoRA) ======================
115
  bnb_config = BitsAndBytesConfig(
116
  load_in_4bit=True,
 
124
  quantization_config=bnb_config,
125
  device_map="auto", # automatically puts layers on GPU
126
  trust_remote_code=True,
127
+ **_from_pretrained_kwargs_with_token(hf_token),
128
  )
129
 
130
+ tokenizer = AutoTokenizer.from_pretrained(
131
+ model_name,
132
+ trust_remote_code=True,
133
+ **_from_pretrained_kwargs_with_token(hf_token),
134
+ )
135
  tokenizer.pad_token = tokenizer.eos_token
136
 
137
  # Prepare for QLoRA
 
151
  model.print_trainable_parameters() # should show ~0.5-1% trainable params
152
 
153
  # ====================== LOAD & PREPARE DATASET ======================
154
+ try:
155
+ dataset = load_dataset(dataset_name, split="train", token=hf_token) if hf_token else load_dataset(dataset_name, split="train")
156
+ except TypeError:
157
+ # Fallback for older datasets APIs.
158
+ dataset = load_dataset(dataset_name, split="train", use_auth_token=hf_token) if hf_token else load_dataset(dataset_name, split="train")
159
 
160
  # Convert ShareGPT "conversations" β†’ standard messages format
161
  def map_to_messages(example):
 
221
  print(f"βœ… Training finished! LoRA adapter saved to {output_dir}")
222
 
223
 
224
+ merged_dir = "llama32-1b-finetuned-merged"
 
 
 
 
 
 
 
 
 
 
 
 
 
225
  model = AutoPeftModelForCausalLM.from_pretrained(
226
  output_dir,
227
  device_map="auto",
228
+ torch_dtype=torch.bfloat16,
229
  )
230
  model = model.merge_and_unload()
231
+ model.save_pretrained(merged_dir)
232
+ tokenizer.save_pretrained(merged_dir)
233
+ print(f"βœ… Merged model saved to {merged_dir}")