joshua400 commited on
Commit
5cb03a6
·
1 Parent(s): 8084d88

Fix: Added PeftModel loading logic for LoRA adapters and updated requirements

Browse files
Files changed (2) hide show
  1. inference.py +39 -1
  2. requirements.txt +2 -0
inference.py CHANGED
@@ -100,10 +100,48 @@ class TrainedInferencePolicy:
100
  def __init__(self, model_name: str = "Joshua1702/fairrecovery-Qwen2.5-7B-GRPO"):
101
  import torch
102
  from transformers import AutoModelForCausalLM, AutoTokenizer
 
 
 
 
 
103
  print(f"Loading model: {model_name}")
104
  self.tokenizer = AutoTokenizer.from_pretrained(model_name)
105
  dtype = torch.float16 if torch.cuda.is_available() else torch.float32
106
- self.model = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=dtype, device_map="auto")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
107
  self.model.eval()
108
 
109
  def __call__(self, obs: FairRecoveryObservation) -> FairRecoveryAction:
 
100
  def __init__(self, model_name: str = "Joshua1702/fairrecovery-Qwen2.5-7B-GRPO"):
101
  import torch
102
  from transformers import AutoModelForCausalLM, AutoTokenizer
103
+ try:
104
+ from peft import PeftModel
105
+ except ImportError:
106
+ PeftModel = None
107
+
108
  print(f"Loading model: {model_name}")
109
  self.tokenizer = AutoTokenizer.from_pretrained(model_name)
110
  dtype = torch.float16 if torch.cuda.is_available() else torch.float32
111
+
112
+ # Hardcoded mapping for known adapters to their base models
113
+ BASE_MODELS = {
114
+ "Joshua1702/fairrecovery-Qwen2.5-7B-GRPO": "unsloth/Qwen2.5-7B-Instruct-bnb-4bit",
115
+ "Joshua1702/fairrecovery-llama-1b-grpo": "unsloth/Llama-3.2-1B-Instruct-bnb-4bit",
116
+ "Joshua1702/fairrecovery-Llama-3.2-1B": "unsloth/Llama-3.2-1B-Instruct-bnb-4bit"
117
+ }
118
+
119
+ base_model_id = BASE_MODELS.get(model_name)
120
+
121
+ try:
122
+ if base_model_id and PeftModel:
123
+ print(f"Detected adapter. Loading base model: {base_model_id}")
124
+ base_model = AutoModelForCausalLM.from_pretrained(
125
+ base_model_id,
126
+ torch_dtype=dtype,
127
+ device_map="auto"
128
+ )
129
+ self.model = PeftModel.from_pretrained(base_model, model_name)
130
+ else:
131
+ self.model = AutoModelForCausalLM.from_pretrained(
132
+ model_name,
133
+ torch_dtype=dtype,
134
+ device_map="auto"
135
+ )
136
+ except Exception as e:
137
+ print(f"Standard load failed: {e}. Trying fallback...")
138
+ # If standard load fails, it might be because it's an adapter but not in our mapping
139
+ self.model = AutoModelForCausalLM.from_pretrained(
140
+ model_name,
141
+ torch_dtype=dtype,
142
+ device_map="auto"
143
+ )
144
+
145
  self.model.eval()
146
 
147
  def __call__(self, obs: FairRecoveryObservation) -> FairRecoveryAction:
requirements.txt CHANGED
@@ -13,6 +13,8 @@ huggingface_hub
13
  transformers
14
  torch
15
  accelerate
 
 
16
  sentencepiece
17
  pandas
18
  matplotlib
 
13
  transformers
14
  torch
15
  accelerate
16
+ peft
17
+ bitsandbytes
18
  sentencepiece
19
  pandas
20
  matplotlib