PrithviRana/DevOps
2101
1import os2import sys3import torch4 5from transformers import (6 AutoTokenizer,7 AutoModelForCausalLM,8)9 10from peft import PeftModel11 12 13# ============================================================14# CONFIGURATION15# ============================================================16 17BASE_MODEL = "Qwen/Qwen2.5-3B-Instruct"18 19LORA_PATH = "/root/ai-tuning/lora-output"20 21MAX_LENGTH = 25622 23MAX_NEW_TOKENS = 15024 25MIN_NEW_TOKENS = 1026 27 28# ============================================================29# CPU CONFIGURATION30# ============================================================31 32CPU_COUNT = os.cpu_count() or 433 34torch.set_num_threads(CPU_COUNT)35 36torch.set_num_interop_threads(2)37 38 39# ============================================================40# SYSTEM INFORMATION41# ============================================================42 43print()44print("==========================================")45print("SYSTEM INFORMATION")46print("==========================================")47 48print("CPU threads :", CPU_COUNT)49print("PyTorch threads :", torch.get_num_threads())50print("CUDA available :", torch.cuda.is_available())51print("PyTorch version :", torch.__version__)52print("Base model :", BASE_MODEL)53print("LoRA adapter :", LORA_PATH)54 55 56# ============================================================57# CHECK LoRA DIRECTORY58# ============================================================59 60print()61print("==========================================")62print("CHECKING LoRA ADAPTER")63print("==========================================")64 65if not os.path.isdir(LORA_PATH):66 67 print("ERROR: LoRA directory not found:")68 print(LORA_PATH)69 70 sys.exit(1)71 72 73adapter_file = os.path.join(74 LORA_PATH,75 "adapter_model.safetensors"76)77 78adapter_file_bin = os.path.join(79 LORA_PATH,80 "adapter_model.bin"81)82 83 84if not os.path.exists(adapter_file) and not os.path.exists(adapter_file_bin):85 86 print("ERROR: LoRA adapter file not found.")87 88 print()89 print("Expected:")90 print(adapter_file)91 print("OR")92 print(adapter_file_bin)93 94 print()95 print("Files found:")96 97 for filename in sorted(os.listdir(LORA_PATH)):98 print(" ", filename)99 100 sys.exit(1)101 102 103print("LoRA adapter found")104 105 106# ============================================================107# LOAD TOKENIZER108# ============================================================109 110print()111print("==========================================")112print("LOADING TOKENIZER")113print("==========================================")114 115try:116 117 tokenizer = AutoTokenizer.from_pretrained(118 LORA_PATH,119 use_fast=True,120 )121 122except Exception as error:123 124 print("ERROR loading tokenizer:")125 print(type(error).__name__)126 print(error)127 128 sys.exit(1)129 130 131# ------------------------------------------------------------132# PAD TOKEN133# ------------------------------------------------------------134 135if tokenizer.pad_token is None:136 137 tokenizer.pad_token = tokenizer.eos_token138 139 140print("Tokenizer loaded successfully")141 142print()143print("TOKENIZER INFORMATION")144print("------------------------------------------")145 146print("EOS token :", repr(tokenizer.eos_token))147print("EOS token ID :", tokenizer.eos_token_id)148 149print("PAD token :", repr(tokenizer.pad_token))150print("PAD token ID :", tokenizer.pad_token_id)151 152print("BOS token :", repr(tokenizer.bos_token))153print("BOS token ID :", tokenizer.bos_token_id)154 155 156# ============================================================157# LOAD BASE MODEL158# ============================================================159 160print()161print("==========================================")162print("LOADING QWEN2.5-3B-INSTRUCT")163print("==========================================")164 165try:166 167 base_model = AutoModelForCausalLM.from_pretrained(168 BASE_MODEL,169 torch_dtype=torch.float32,170 )171 172except Exception as error:173 174 print("ERROR loading base model:")175 print(type(error).__name__)176 print(error)177 178 sys.exit(1)179 180 181# ------------------------------------------------------------182# Configure padding183# ------------------------------------------------------------184 185base_model.config.pad_token_id = tokenizer.pad_token_id186 187print("Base model loaded successfully")188 189 190# ============================================================191# LOAD LoRA ADAPTER192# ============================================================193 194print()195print("==========================================")196print("LOADING LoRA ADAPTER")197print("==========================================")198 199try:200 201 model = PeftModel.from_pretrained(202 base_model,203 LORA_PATH,204 )205 206except Exception as error:207 208 print("ERROR loading LoRA adapter:")209 print(type(error).__name__)210 print(error)211 212 sys.exit(1)213 214 215# ------------------------------------------------------------216# Evaluation mode217# ------------------------------------------------------------218 219model.eval()220 221 222print("LoRA adapter loaded successfully")223 224 225# ============================================================226# MODEL INFORMATION227# ============================================================228 229print()230print("==========================================")231print("MODEL INFORMATION")232print("==========================================")233 234model.print_trainable_parameters()235 236 237# ============================================================238# GENERATION CONFIGURATION239# ============================================================240 241print()242print("==========================================")243print("GENERATION CONFIGURATION")244print("==========================================")245 246# ------------------------------------------------------------247# Deterministic generation248#249# do_sample=False means:250#251# temperature = not used252# top_p = not used253# top_k = not used254#255# This removes the warnings you were seeing.256# ------------------------------------------------------------257 258model.generation_config.do_sample = False259 260model.generation_config.temperature = None261 262model.generation_config.top_p = None263 264model.generation_config.top_k = None265 266 267print("do_sample :", model.generation_config.do_sample)268 269print("temperature :", model.generation_config.temperature)270 271print("top_p :", model.generation_config.top_p)272 273print("top_k :", model.generation_config.top_k)274 275print("repetition_penalty :", 1.1)276 277print("max_new_tokens :", MAX_NEW_TOKENS)278 279print("min_new_tokens :", MIN_NEW_TOKENS)280 281 282# ============================================================283# GENERATION FUNCTION284# ============================================================285 286def ask_devops(question):287 288 # --------------------------------------------------------289 # IMPORTANT290 #291 # This format matches the format used during training.292 #293 # Training:294 #295 # ### Instruction:296 # question297 #298 # ### Input:299 #300 # ### Response:301 # answer302 #303 # --------------------------------------------------------304 305 prompt = (306 "### Instruction:\n"307 f"{question}\n\n"308 "### Input:\n"309 "\n"310 "### Response:\n"311 )312 313 314 # --------------------------------------------------------315 # PRINT PROMPT316 # --------------------------------------------------------317 318 print()319 print("Prompt:")320 print("------------------------------------------")321 print(prompt)322 print("------------------------------------------")323 324 325 # --------------------------------------------------------326 # TOKENIZE327 # --------------------------------------------------------328 329 try:330 331 inputs = tokenizer(332 prompt,333 return_tensors="pt",334 truncation=True,335 max_length=MAX_LENGTH,336 padding=False,337 )338 339 except Exception as error:340 341 print()342 print("TOKENIZATION ERROR:")343 print(type(error).__name__)344 print(error)345 346 return ""347 348 349 # --------------------------------------------------------350 # DEBUG INPUT351 # --------------------------------------------------------352 353 input_token_count = inputs["input_ids"].shape[1]354 355 print()356 print("DEBUG INPUT")357 print("------------------------------------------")358 359 print("Input token count :", input_token_count)360 361 print("EOS token ID :", tokenizer.eos_token_id)362 363 print("PAD token ID :", tokenizer.pad_token_id)364 365 366 # --------------------------------------------------------367 # GENERATE368 # --------------------------------------------------------369 370 try:371 372 with torch.inference_mode():373 374 outputs = model.generate(375 376 input_ids=inputs["input_ids"],377 378 attention_mask=inputs["attention_mask"],379 380 381 # ------------------------------------------------382 # Generation length383 # ------------------------------------------------384 385 max_new_tokens=MAX_NEW_TOKENS,386 387 min_new_tokens=MIN_NEW_TOKENS,388 389 390 # ------------------------------------------------391 # Deterministic generation392 # ------------------------------------------------393 394 do_sample=False,395 396 397 # ------------------------------------------------398 # Repetition control399 # ------------------------------------------------400 401 repetition_penalty=1.1,402 403 404 # ------------------------------------------------405 # Tokens406 # ------------------------------------------------407 408 pad_token_id=tokenizer.pad_token_id,409 410 eos_token_id=tokenizer.eos_token_id,411 412 413 # ------------------------------------------------414 # KV cache415 # ------------------------------------------------416 417 use_cache=True,418 )419 420 except Exception as error:421 422 print()423 print("GENERATION ERROR:")424 print(type(error).__name__)425 print(error)426 427 return ""428 429 430 # ========================================================431 # REMOVE INPUT PROMPT432 # ========================================================433 434 input_length = inputs["input_ids"].shape[1]435 436 generated_tokens = outputs[437 0,438 input_length:439 ]440 441 442 # ========================================================443 # DEBUG GENERATED TOKENS444 # ========================================================445 446 print()447 print("DEBUG OUTPUT")448 print("------------------------------------------")449 450 print(451 "Generated token count :",452 len(generated_tokens)453 )454 455 print(456 "Generated token IDs :",457 generated_tokens[:30].tolist()458 )459 460 461 # --------------------------------------------------------462 # Decode463 # --------------------------------------------------------464 465 answer = tokenizer.decode(466 generated_tokens,467 skip_special_tokens=True,468 )469 470 471 # --------------------------------------------------------472 # Clean answer473 # --------------------------------------------------------474 475 answer = answer.strip()476 477 478 # ========================================================479 # RETURN480 # ========================================================481 482 return answer483 484 485# ============================================================486# TEST QUESTIONS487# ============================================================488 489questions = [490 491 "How do I check a Linux server's uptime?",492 493 "How do I check memory usage in Linux?",494 495 "How do I check CPU usage in Linux?",496 497 "How do I check running Docker containers?",498 499 "How do I restart a Docker container?",500 501 "How do I check nginx error logs?",502 503 "How do I troubleshoot HTTP 502 error in nginx?",504 505 "How do I check disk space used by a directory?",506 507]508 509 510# ============================================================511# AUTOMATIC TEST512# ============================================================513 514print()515print("==========================================")516print("STARTING LoRA MODEL TEST")517print("==========================================")518 519print()520print("Number of test questions :", len(questions))521 522print()523print("IMPORTANT:")524print("The first question may take some time on CPU.")525print("Please wait for the generated answer.")526 527 528for number, question in enumerate(529 questions,530 start=1531):532 533 print()534 print()535 print("##########################################")536 print(f"TEST {number}")537 print("##########################################")538 539 540 print()541 print("Question:")542 print(question)543 544 545 print()546 print("Answer:")547 548 549 try:550 551 answer = ask_devops(question)552 553 554 if answer:555 556 print()557 print("==========================================")558 print("MODEL ANSWER")559 print("==========================================")560 print(answer)561 562 else:563 564 print()565 print("[EMPTY RESPONSE]")566 567 568 except Exception as error:569 570 print()571 print("ERROR:")572 print(type(error).__name__)573 print(error)574 575 576# ============================================================577# INTERACTIVE MODE578# ============================================================579 580print()581print()582print("==========================================")583print("INTERACTIVE DEVOPS CHAT")584print("==========================================")585 586print()587print("Enter your DevOps question.")588 589print("Type 'exit' to stop.")590 591print()592 593 594while True:595 596 try:597 598 question = input("\nYou: ").strip()599 600 601 except KeyboardInterrupt:602 603 print()604 print()605 print("Exiting...")606 break607 608 609 except EOFError:610 611 print()612 print()613 print("Exiting...")614 break615 616 617 # --------------------------------------------------------618 # EXIT619 # --------------------------------------------------------620 621 if question.lower() in [622 "exit",623 "quit",624 "q",625 ]:626 627 print()628 print("Exiting...")629 break630 631 632 # --------------------------------------------------------633 # EMPTY INPUT634 # --------------------------------------------------------635 636 if not question:637 638 continue639 640 641 # --------------------------------------------------------642 # GENERATE ANSWER643 # --------------------------------------------------------644 645 print()646 print("AI:")647 648 649 try:650 651 answer = ask_devops(question)652 653 654 if answer:655 656 print()657 print(answer)658 659 else:660 661 print()662 print("[EMPTY RESPONSE]")663 664 665 except Exception as error:666 667 print()668 print("ERROR:")669 print(type(error).__name__)670 print(error)671 672 673# ============================================================674# COMPLETE675# ============================================================676 677print()678print("==========================================")679print("LoRA TEST COMPLETE")680print("==========================================")681 