Team Ai
Modelpublic

PrithviRana/DevOps

sourceHugging Faceupdated 1mo agoView on Hugging Face
2likes101downloads
test_lora.py681 linesDownload Raw Back to root
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