DevOps / test_lora.py
PrithviRana's picture
Upload folder using huggingface_hub
9a40336 verified
Raw History Blame Contribute Delete
15 kB
import os
import sys
import torch
from transformers import (
AutoTokenizer,
AutoModelForCausalLM,
)
from peft import PeftModel
# ============================================================
# CONFIGURATION
# ============================================================
BASE_MODEL = "Qwen/Qwen2.5-3B-Instruct"
LORA_PATH = "/root/ai-tuning/lora-output"
MAX_LENGTH = 256
MAX_NEW_TOKENS = 150
MIN_NEW_TOKENS = 10
# ============================================================
# CPU CONFIGURATION
# ============================================================
CPU_COUNT = os.cpu_count() or 4
torch.set_num_threads(CPU_COUNT)
torch.set_num_interop_threads(2)
# ============================================================
# SYSTEM INFORMATION
# ============================================================
print()
print("==========================================")
print("SYSTEM INFORMATION")
print("==========================================")
print("CPU threads :", CPU_COUNT)
print("PyTorch threads :", torch.get_num_threads())
print("CUDA available :", torch.cuda.is_available())
print("PyTorch version :", torch.__version__)
print("Base model :", BASE_MODEL)
print("LoRA adapter :", LORA_PATH)
# ============================================================
# CHECK LoRA DIRECTORY
# ============================================================
print()
print("==========================================")
print("CHECKING LoRA ADAPTER")
print("==========================================")
if not os.path.isdir(LORA_PATH):
print("ERROR: LoRA directory not found:")
print(LORA_PATH)
sys.exit(1)
adapter_file = os.path.join(
LORA_PATH,
"adapter_model.safetensors"
)
adapter_file_bin = os.path.join(
LORA_PATH,
"adapter_model.bin"
)
if not os.path.exists(adapter_file) and not os.path.exists(adapter_file_bin):
print("ERROR: LoRA adapter file not found.")
print()
print("Expected:")
print(adapter_file)
print("OR")
print(adapter_file_bin)
print()
print("Files found:")
for filename in sorted(os.listdir(LORA_PATH)):
print(" ", filename)
sys.exit(1)
print("LoRA adapter found")
# ============================================================
# LOAD TOKENIZER
# ============================================================
print()
print("==========================================")
print("LOADING TOKENIZER")
print("==========================================")
try:
tokenizer = AutoTokenizer.from_pretrained(
LORA_PATH,
use_fast=True,
)
except Exception as error:
print("ERROR loading tokenizer:")
print(type(error).__name__)
print(error)
sys.exit(1)
# ------------------------------------------------------------
# PAD TOKEN
# ------------------------------------------------------------
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
print("Tokenizer loaded successfully")
print()
print("TOKENIZER INFORMATION")
print("------------------------------------------")
print("EOS token :", repr(tokenizer.eos_token))
print("EOS token ID :", tokenizer.eos_token_id)
print("PAD token :", repr(tokenizer.pad_token))
print("PAD token ID :", tokenizer.pad_token_id)
print("BOS token :", repr(tokenizer.bos_token))
print("BOS token ID :", tokenizer.bos_token_id)
# ============================================================
# LOAD BASE MODEL
# ============================================================
print()
print("==========================================")
print("LOADING QWEN2.5-3B-INSTRUCT")
print("==========================================")
try:
base_model = AutoModelForCausalLM.from_pretrained(
BASE_MODEL,
torch_dtype=torch.float32,
)
except Exception as error:
print("ERROR loading base model:")
print(type(error).__name__)
print(error)
sys.exit(1)
# ------------------------------------------------------------
# Configure padding
# ------------------------------------------------------------
base_model.config.pad_token_id = tokenizer.pad_token_id
print("Base model loaded successfully")
# ============================================================
# LOAD LoRA ADAPTER
# ============================================================
print()
print("==========================================")
print("LOADING LoRA ADAPTER")
print("==========================================")
try:
model = PeftModel.from_pretrained(
base_model,
LORA_PATH,
)
except Exception as error:
print("ERROR loading LoRA adapter:")
print(type(error).__name__)
print(error)
sys.exit(1)
# ------------------------------------------------------------
# Evaluation mode
# ------------------------------------------------------------
model.eval()
print("LoRA adapter loaded successfully")
# ============================================================
# MODEL INFORMATION
# ============================================================
print()
print("==========================================")
print("MODEL INFORMATION")
print("==========================================")
model.print_trainable_parameters()
# ============================================================
# GENERATION CONFIGURATION
# ============================================================
print()
print("==========================================")
print("GENERATION CONFIGURATION")
print("==========================================")
# ------------------------------------------------------------
# Deterministic generation
#
# do_sample=False means:
#
# temperature = not used
# top_p = not used
# top_k = not used
#
# This removes the warnings you were seeing.
# ------------------------------------------------------------
model.generation_config.do_sample = False
model.generation_config.temperature = None
model.generation_config.top_p = None
model.generation_config.top_k = None
print("do_sample :", model.generation_config.do_sample)
print("temperature :", model.generation_config.temperature)
print("top_p :", model.generation_config.top_p)
print("top_k :", model.generation_config.top_k)
print("repetition_penalty :", 1.1)
print("max_new_tokens :", MAX_NEW_TOKENS)
print("min_new_tokens :", MIN_NEW_TOKENS)
# ============================================================
# GENERATION FUNCTION
# ============================================================
def ask_devops(question):
# --------------------------------------------------------
# IMPORTANT
#
# This format matches the format used during training.
#
# Training:
#
# ### Instruction:
# question
#
# ### Input:
#
# ### Response:
# answer
#
# --------------------------------------------------------
prompt = (
"### Instruction:\n"
f"{question}\n\n"
"### Input:\n"
"\n"
"### Response:\n"
)
# --------------------------------------------------------
# PRINT PROMPT
# --------------------------------------------------------
print()
print("Prompt:")
print("------------------------------------------")
print(prompt)
print("------------------------------------------")
# --------------------------------------------------------
# TOKENIZE
# --------------------------------------------------------
try:
inputs = tokenizer(
prompt,
return_tensors="pt",
truncation=True,
max_length=MAX_LENGTH,
padding=False,
)
except Exception as error:
print()
print("TOKENIZATION ERROR:")
print(type(error).__name__)
print(error)
return ""
# --------------------------------------------------------
# DEBUG INPUT
# --------------------------------------------------------
input_token_count = inputs["input_ids"].shape[1]
print()
print("DEBUG INPUT")
print("------------------------------------------")
print("Input token count :", input_token_count)
print("EOS token ID :", tokenizer.eos_token_id)
print("PAD token ID :", tokenizer.pad_token_id)
# --------------------------------------------------------
# GENERATE
# --------------------------------------------------------
try:
with torch.inference_mode():
outputs = model.generate(
input_ids=inputs["input_ids"],
attention_mask=inputs["attention_mask"],
# ------------------------------------------------
# Generation length
# ------------------------------------------------
max_new_tokens=MAX_NEW_TOKENS,
min_new_tokens=MIN_NEW_TOKENS,
# ------------------------------------------------
# Deterministic generation
# ------------------------------------------------
do_sample=False,
# ------------------------------------------------
# Repetition control
# ------------------------------------------------
repetition_penalty=1.1,
# ------------------------------------------------
# Tokens
# ------------------------------------------------
pad_token_id=tokenizer.pad_token_id,
eos_token_id=tokenizer.eos_token_id,
# ------------------------------------------------
# KV cache
# ------------------------------------------------
use_cache=True,
)
except Exception as error:
print()
print("GENERATION ERROR:")
print(type(error).__name__)
print(error)
return ""
# ========================================================
# REMOVE INPUT PROMPT
# ========================================================
input_length = inputs["input_ids"].shape[1]
generated_tokens = outputs[
0,
input_length:
]
# ========================================================
# DEBUG GENERATED TOKENS
# ========================================================
print()
print("DEBUG OUTPUT")
print("------------------------------------------")
print(
"Generated token count :",
len(generated_tokens)
)
print(
"Generated token IDs :",
generated_tokens[:30].tolist()
)
# --------------------------------------------------------
# Decode
# --------------------------------------------------------
answer = tokenizer.decode(
generated_tokens,
skip_special_tokens=True,
)
# --------------------------------------------------------
# Clean answer
# --------------------------------------------------------
answer = answer.strip()
# ========================================================
# RETURN
# ========================================================
return answer
# ============================================================
# TEST QUESTIONS
# ============================================================
questions = [
"How do I check a Linux server's uptime?",
"How do I check memory usage in Linux?",
"How do I check CPU usage in Linux?",
"How do I check running Docker containers?",
"How do I restart a Docker container?",
"How do I check nginx error logs?",
"How do I troubleshoot HTTP 502 error in nginx?",
"How do I check disk space used by a directory?",
]
# ============================================================
# AUTOMATIC TEST
# ============================================================
print()
print("==========================================")
print("STARTING LoRA MODEL TEST")
print("==========================================")
print()
print("Number of test questions :", len(questions))
print()
print("IMPORTANT:")
print("The first question may take some time on CPU.")
print("Please wait for the generated answer.")
for number, question in enumerate(
questions,
start=1
):
print()
print()
print("##########################################")
print(f"TEST {number}")
print("##########################################")
print()
print("Question:")
print(question)
print()
print("Answer:")
try:
answer = ask_devops(question)
if answer:
print()
print("==========================================")
print("MODEL ANSWER")
print("==========================================")
print(answer)
else:
print()
print("[EMPTY RESPONSE]")
except Exception as error:
print()
print("ERROR:")
print(type(error).__name__)
print(error)
# ============================================================
# INTERACTIVE MODE
# ============================================================
print()
print()
print("==========================================")
print("INTERACTIVE DEVOPS CHAT")
print("==========================================")
print()
print("Enter your DevOps question.")
print("Type 'exit' to stop.")
print()
while True:
try:
question = input("\nYou: ").strip()
except KeyboardInterrupt:
print()
print()
print("Exiting...")
break
except EOFError:
print()
print()
print("Exiting...")
break
# --------------------------------------------------------
# EXIT
# --------------------------------------------------------
if question.lower() in [
"exit",
"quit",
"q",
]:
print()
print("Exiting...")
break
# --------------------------------------------------------
# EMPTY INPUT
# --------------------------------------------------------
if not question:
continue
# --------------------------------------------------------
# GENERATE ANSWER
# --------------------------------------------------------
print()
print("AI:")
try:
answer = ask_devops(question)
if answer:
print()
print(answer)
else:
print()
print("[EMPTY RESPONSE]")
except Exception as error:
print()
print("ERROR:")
print(type(error).__name__)
print(error)
# ============================================================
# COMPLETE
# ============================================================
print()
print("==========================================")
print("LoRA TEST COMPLETE")
print("==========================================")