Download llm_handler.py from kprsh/DRLLMTRY: direct link, hf CLI and curl.
- Browser
- Download file 6.46 kB
-
https://huggingface.co/spaces/kprsh/DRLLMTRY/resolve/main/llm_handler.py
- Command line
-
hf download hf://spaces/kprsh/DRLLMTRY/llm_handler.py
-
curl -L -o llm_handler.py https://huggingface.co/spaces/kprsh/DRLLMTRY/resolve/main/llm_handler.py
6.46 kB
| import pandas as pd | |
| import re | |
| import os | |
| from pathlib import Path | |
| from huggingface_hub import InferenceClient | |
| from dotenv import load_dotenv | |
| class DDoSInference: | |
| def __init__(self): | |
| """ | |
| Initialize DDoSInference class, set up the API client, and paths for dataset and results. | |
| """ | |
| load_dotenv() | |
| self.client = InferenceClient(api_key=os.getenv("HF_TOK_KEY")) | |
| self.model = "Qwen/Qwen2.5-Coder-32B-Instruct" | |
| self.dataset_path = Path("~/.dataset/original.csv").expanduser() | |
| self.results_path = Path("~/.dataset/PROBABILITY_OF_EACH_ROW_DDOS_AND_BENIGN.csv").expanduser() | |
| self.results_path.parent.mkdir(parents=True, exist_ok=True) | |
| def process_dataset(self): | |
| """ | |
| Process the dataset row by row, performing inference using the LLM for each row. | |
| """ | |
| if not self.dataset_path.exists(): | |
| raise FileNotFoundError("The preprocessed dataset file does not exist. Ensure it is generated using the processor.") | |
| ddos_data = pd.read_csv(self.dataset_path) | |
| label_column = " Label" | |
| if label_column not in ddos_data.columns: | |
| label_column = input("Enter the label column name in your dataset: ").strip() | |
| if label_column not in ddos_data.columns: | |
| raise ValueError(f"Label column '{label_column}' not found in the dataset.") | |
| ddos_data_without_label = ddos_data.drop([label_column], axis=1) | |
| stats = { | |
| 'Max': ddos_data_without_label.max(), | |
| 'Min': ddos_data_without_label.min(), | |
| 'Median': ddos_data_without_label.median(), | |
| 'Mean': ddos_data_without_label.mean(), | |
| 'Variance': ddos_data_without_label.var() | |
| } | |
| # Generate knowledge prompt | |
| know_prompt = self.generate_knowledge_prompt(stats) | |
| # Prepare results DataFrame | |
| predict_df = self.load_or_create_results() | |
| start_index = predict_df.shape[0] | |
| print(f"Starting inference from row {start_index}") | |
| # Process each row for inference | |
| for i in range(start_index, ddos_data.shape[0]): | |
| row_prompt = self.generate_row_prompt(ddos_data.iloc[i]) | |
| probabilities = self.infer_row(know_prompt, row_prompt) | |
| # If no valid response, mark as "None" | |
| predict_df.loc[i] = [i, *probabilities] if probabilities else [i, "None", "None", "No valid response"] | |
| # Save after each row for resilience | |
| predict_df.to_csv(self.results_path, index=False) | |
| print(f"Processed row {i}: {predict_df.loc[i].to_dict()}") | |
| print("Inference complete. Results saved at:", self.results_path) | |
| def generate_knowledge_prompt(self, stats): | |
| """ | |
| Generates the knowledge prompt based on dataset statistics. | |
| """ | |
| prompt = ( | |
| "Supposed that you are now an [[ HIGHLY EXPERIENCED NETWORK TRAFFIC DATA ANALYSIS EXPERT ]]. " | |
| "You need to help me analyze the data in the DDoS dataset and determine whether the data is [[ DDoS traffic ]] or [[ normal traffic ]]. " | |
| "Here are the maximum, minimum, median, mean, and variance of each column in the dataset to help your judgment:\n" | |
| ) | |
| for col, values in stats.items(): | |
| prompt += f"{col}: max={values:.2f}, min={values:.2f}, median={values:.2f}, mean={values:.2f}, variance={values:.2f}\n" | |
| return prompt | |
| def generate_row_prompt(self, row): | |
| """ | |
| Generates a row-specific prompt for the LLM. | |
| """ | |
| row_prompt = ( | |
| "Next, I will give you a piece of data about network traffic information. " | |
| "You need to tell me the probability of this data being DDoS traffic or normal traffic. " | |
| "Express the probability in the format [0.xxx, 0.xxx], where the first number represents DDoS probability and the second represents normal traffic probability. " | |
| "Ensure that the sum of probabilities is exactly 1.\n" | |
| ) | |
| for col, val in row.items(): | |
| row_prompt += f"{col}: {val}, " | |
| return row_prompt.strip(', ') | |
| def infer_row(self, know_prompt, row_prompt): | |
| """ | |
| Performs inference for a single row using the LLM. | |
| """ | |
| try: | |
| messages = [ | |
| {'role': 'user', 'content': know_prompt}, | |
| {'role': 'user', 'content': row_prompt} | |
| ] | |
| completion = self.client.chat.completions.create( | |
| model=self.model, | |
| messages=messages, | |
| max_tokens=1000 | |
| ) | |
| response = completion.choices[0].message.content | |
| probabilities = self.extract_probabilities(response) | |
| return probabilities | |
| except Exception as e: | |
| print(f"Error during inference for row: {e}") | |
| return None | |
| def extract_probabilities(self, response): | |
| """ | |
| Extract probabilities from the LLM response using regex. | |
| """ | |
| pattern = r'\[(.*?)\]' | |
| match = re.search(pattern, response) | |
| if match: | |
| probs = match.group(1).split(',') | |
| return [float(p.strip()) for p in probs if p.strip()] | |
| return None | |
| def get_chat_response(self, user_input): | |
| """ | |
| Generate a response for the user's question using the LLM. | |
| """ | |
| try: | |
| messages = [{'role': 'user', 'content': user_input}] | |
| completion = self.client.chat.completions.create( | |
| model=self.model, | |
| messages=messages, | |
| max_tokens=500 | |
| ) | |
| response = completion.choices[0].message.content | |
| return response.strip() | |
| except Exception as e: | |
| return f"Error: Unable to process your request due to {e}." | |
| def load_or_create_results(self): | |
| """ | |
| Loads the existing results or creates a new DataFrame if the results file doesn't exist. | |
| """ | |
| if self.results_path.exists(): | |
| return pd.read_csv(self.results_path) | |
| else: | |
| return pd.DataFrame(columns=["index", "attack", "benign", "original"]) | |
| # Example usage | |
| if __name__ == "__main__": | |
| handler = DDoSInference() | |
| handler.process_dataset() | |
| print("You can now interact with the model for mitigation steps or download the results.") | |