xxccho commited on
Commit
42b0aaf
ยท
verified ยท
1 Parent(s): b375100

Update README.md

Browse files
Files changed (1) hide show
  1. README.md +33 -5
README.md CHANGED
@@ -23,11 +23,35 @@ import torch
23
  from transformers import AutoTokenizer, AutoModelForSequenceClassification
24
  from peft import PeftModel, PeftConfig
25
 
26
- # 1. Define the PEFT model ID
 
 
27
  peft_model_id = "xxccho/margin_reg_baseline"
28
 
29
- # 2. Load the PEFT config
30
- config = PeftConfig.from_pretrained(peft_model_id)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31
 
32
  # 3. Load tokenizer from base model (safer)
33
  tokenizer = AutoTokenizer.from_pretrained(config.base_model_name_or_path)
@@ -44,8 +68,12 @@ base_model = AutoModelForSequenceClassification.from_pretrained(
44
  device_map="auto"
45
  )
46
 
47
- # 5. Apply LoRA adapter
48
- model = PeftModel.from_pretrained(base_model, peft_model_id)
 
 
 
 
49
  model.config.pad_token_id = tokenizer.pad_token_id
50
  model.eval()
51
 
 
23
  from transformers import AutoTokenizer, AutoModelForSequenceClassification
24
  from peft import PeftModel, PeftConfig
25
 
26
+ # -----------------------------
27
+ # 1. Define the PEFT model ID & Checkpoint (Epoch)
28
+ # -----------------------------
29
  peft_model_id = "xxccho/margin_reg_baseline"
30
 
31
+ # [Optional] ํŠน์ • Epoch์˜ ์ค‘๊ฐ„ ์ฒดํฌํฌ์ธํŠธ๋ฅผ ๋ถˆ๋Ÿฌ์˜ค๊ณ  ์‹ถ์„ ๋•Œ ์•„๋ž˜ ๋ณ€์ˆ˜๋ฅผ ์ง€์ •ํ•˜์„ธ์š”.
32
+ # ์ง€์ •ํ•˜์ง€ ์•Š๊ณ  None์œผ๋กœ ๋‘๋ฉด ๋ ˆํฌ์ง€ํ† ๋ฆฌ ์ตœ์ƒ๋‹จ์— ์žˆ๋Š” ๋งˆ์ง€๋ง‰(์ตœ์ข…) ํ•™์Šต ๋ชจ๋ธ์ด ๋กœ๋“œ๋ฉ๋‹ˆ๋‹ค.
33
+ #
34
+ # [ Ckeckpoints to Epochs Mapping ]
35
+ # Epoch 1 : "checkpoint-246"
36
+ # Epoch 2 : "checkpoint-492"
37
+ # Epoch 3 : "checkpoint-738"
38
+ # Epoch 4 : "checkpoint-984"
39
+ # Epoch 5 : "checkpoint-1230"
40
+ # Epoch 6 : "checkpoint-1476"
41
+ # Epoch 7 : "checkpoint-1722"
42
+ # Epoch 8 : "checkpoint-1968"
43
+ # Epoch 9 : "checkpoint-2214"
44
+ # Epoch 10 : "checkpoint-2460"
45
+
46
+ # ์˜ˆ์‹œ: 5 Epoch ์ฒดํฌํฌ์ธํŠธ๋ฅผ ์‚ฌ์šฉํ•˜๋ ค๋ฉด ์•„๋ž˜์™€ ๊ฐ™์ด ๋ณ€๊ฒฝํ•˜์„ธ์š”.
47
+ # checkpoint = "checkpoint-1230"
48
+ checkpoint = None # None์ด๋ฉด ๋””ํดํŠธ๋กœ ์ œ์ผ ๋งˆ์ง€๋ง‰ ์ €์žฅ ๋ชจ๋ธ์„ ์”๋‹ˆ๋‹ค.
49
+
50
+ # 2. Load the PEFT config (์ฒดํฌํฌ์ธํŠธ ์ง€์ • ์—ฌ๋ถ€์— ๋”ฐ๋ผ subfolder ์ ์šฉ)
51
+ if checkpoint:
52
+ config = PeftConfig.from_pretrained(peft_model_id, subfolder=checkpoint)
53
+ else:
54
+ config = PeftConfig.from_pretrained(peft_model_id)
55
 
56
  # 3. Load tokenizer from base model (safer)
57
  tokenizer = AutoTokenizer.from_pretrained(config.base_model_name_or_path)
 
68
  device_map="auto"
69
  )
70
 
71
+ # 5. Apply LoRA adapter (์ฒดํฌํฌ์ธํŠธ ์ง€์ • ์—ฌ๋ถ€์— ๋”ฐ๋ผ subfolder ์ ์šฉ)
72
+ if checkpoint:
73
+ model = PeftModel.from_pretrained(base_model, peft_model_id, subfolder=checkpoint)
74
+ else:
75
+ model = PeftModel.from_pretrained(base_model, peft_model_id)
76
+
77
  model.config.pad_token_id = tokenizer.pad_token_id
78
  model.eval()
79