4moha commited on
Commit
dabebe4
·
verified ·
1 Parent(s): 6766f54

upload train_lora.py (renamed)

Browse files
Files changed (1) hide show
  1. train_lora.py +98 -0
train_lora.py ADDED
@@ -0,0 +1,98 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # /// script
2
+ # requires-python = ">=3.11"
3
+ # dependencies = [
4
+ # "trl>=0.12.0",
5
+ # "peft>=0.7.0",
6
+ # "transformers>=4.45",
7
+ # "datasets>=2.20",
8
+ # "accelerate>=0.34",
9
+ # "trackio",
10
+ # "unsloth",
11
+ # ]
12
+ # ///
13
+ """Phase-A LoRA SFT for the raunch page-mode model — runs inside HF Jobs.
14
+
15
+ Base: Sao10K/L3.1-8B-Stheno-v3.4
16
+ Dataset: 4moha/raunch-page-mode-v0 (private)
17
+ Output: pushed to 4moha/raunch-stheno-v3.4-lora-v0
18
+
19
+ NSFW-only: training data is raunch's NSFW Claude-generated prose. The resulting
20
+ LoRA is deployed to the raunch server instance, NOT the SFW lili server.
21
+
22
+ This script is submitted as the body of the HF Job; it expects the env vars
23
+ HF_TOKEN, HF_DATASET_REPO, HF_MODEL_REPO to be set in the job environment.
24
+ """
25
+ import os
26
+
27
+ from datasets import load_dataset
28
+ from peft import LoraConfig
29
+ from trl import SFTTrainer, SFTConfig
30
+ from unsloth import FastLanguageModel
31
+
32
+
33
+ BASE_MODEL = "Sao10K/L3.1-8B-Stheno-v3.4"
34
+ DATASET_REPO = os.environ.get("HF_DATASET_REPO", "4moha/raunch-page-mode-v0")
35
+ MODEL_REPO = os.environ.get("HF_MODEL_REPO", "4moha/raunch-stheno-v3.4-lora-v0")
36
+
37
+
38
+ def main() -> None:
39
+ # Load model + tokenizer via Unsloth (faster + leaner than vanilla transformers)
40
+ model, tokenizer = FastLanguageModel.from_pretrained(
41
+ model_name=BASE_MODEL,
42
+ max_seq_length=4096,
43
+ dtype=None, # auto
44
+ load_in_4bit=True, # QLoRA — fits more comfortably on A10G
45
+ )
46
+
47
+ model = FastLanguageModel.get_peft_model(
48
+ model,
49
+ r=16,
50
+ lora_alpha=32,
51
+ lora_dropout=0,
52
+ target_modules=["q_proj", "k_proj", "v_proj", "o_proj",
53
+ "gate_proj", "up_proj", "down_proj"],
54
+ use_gradient_checkpointing="unsloth",
55
+ random_state=42,
56
+ )
57
+
58
+ # Load dataset and split off a tiny eval slice for live monitoring during training
59
+ full = load_dataset(DATASET_REPO, data_files="train.jsonl", split="train")
60
+ split = full.train_test_split(test_size=0.05, seed=42)
61
+
62
+ trainer = SFTTrainer(
63
+ model=model,
64
+ tokenizer=tokenizer,
65
+ train_dataset=split["train"],
66
+ eval_dataset=split["test"],
67
+ args=SFTConfig(
68
+ output_dir="raunch-stheno-v3.4-lora-v0",
69
+ push_to_hub=True,
70
+ hub_model_id=MODEL_REPO,
71
+ hub_private_repo=True,
72
+ hub_strategy="every_save",
73
+ num_train_epochs=3,
74
+ per_device_train_batch_size=1,
75
+ gradient_accumulation_steps=8,
76
+ learning_rate=5e-5,
77
+ lr_scheduler_type="cosine",
78
+ warmup_ratio=0.05,
79
+ max_length=4096,
80
+ logging_steps=10,
81
+ save_strategy="steps",
82
+ save_steps=200,
83
+ eval_strategy="steps",
84
+ eval_steps=50,
85
+ seed=42,
86
+ report_to="trackio",
87
+ run_name="raunch-stheno-v3.4-lora-v0",
88
+ project="raunch-page-mode",
89
+ ),
90
+ )
91
+
92
+ trainer.train()
93
+ trainer.push_to_hub()
94
+ print("Training complete. LoRA pushed to:", MODEL_REPO)
95
+
96
+
97
+ if __name__ == "__main__":
98
+ main()