malvec/codebert_Malvec
0
1import os2CACHE_DIR = "/tmp/hf_cache"3os.environ["HF_HOME"] = CACHE_DIR4os.environ["TRANSFORMERS_CACHE"] = CACHE_DIR5os.environ["HUGGINGFACE_HUB_CACHE"] = CACHE_DIR6os.makedirs(CACHE_DIR, exist_ok=True)7 8import torch9from fastapi import FastAPI, UploadFile, File, HTTPException10from fastapi.middleware.cors import CORSMiddleware11from transformers import AutoTokenizer, AutoModelForSequenceClassification12from typing import List, Optional, Dict13from collections import Counter14import json15from huggingface_hub import hf_hub_download16from som_analyzer import SOMAnalyzer 17from datasets import load_dataset 18import numpy as np19 20app = FastAPI()21app.add_middleware(22 CORSMiddleware,23 allow_origins=["*"],24 allow_methods=["*"],25 allow_headers=["*"],26)27 28# ===== 模型設定 =====29MODEL_REPO = "raxhel/Malvec_predict"30HF_TOKEN = os.environ.get("hf_token")31DEVICE = torch.device("cpu")32 33# ===== 載入模型 =====34print("🤗 Loading model...")35tokenizer = AutoTokenizer.from_pretrained(MODEL_REPO, token=HF_TOKEN)36model = AutoModelForSequenceClassification.from_pretrained(37 MODEL_REPO, 38 token=HF_TOKEN,39 output_attentions=True, # ✅ 輸出注意力權重40 output_hidden_states=True # ✅ 關鍵: 輸出 hidden states41)42model.to(DEVICE)43model.eval()44print("✅ Model loaded successfully")45 46# ===== 新增:載入 SOM 模型 =====47print("🧠 Initializing SOM Analyzer...")48SOM_REPO = "malvec/codebert_Malvec" # 49SOM_AVAILABLE = False50try:51 # 52 som_weights_path = hf_hub_download(53 repo_id=SOM_REPO,54 filename="som_weights.npy",55 token=HF_TOKEN # 56 )57 background_data_path = hf_hub_download(58 repo_id=SOM_REPO,59 filename="som_background_data.pkl",60 token=HF_TOKEN # 61 )62 63 som_analyzer = SOMAnalyzer(64 som_weights_path=som_weights_path,65 background_data_path=background_data_path66 )67 print("✅ SOM Analyzer initialized successfully!")68 SOM_AVAILABLE = True69except Exception as e:70 print(f"⚠️ SOM Analyzer initialization failed: {e}")71 print(" SOM analysis will be disabled")72 73 74# ===== 新增:從 main.py 複製過來的輔助函數 =====75# (注意:請從你的 main.py 複製最新的程式碼)76 77async def get_embeddings_from_hf_dataset(filename: str) -> Optional[List[np.ndarray]]:78 # ... (請從 main.py 完整複製貼上) ...79 try:80 print(f"\n📦 Fetching embeddings from HF dataset for: {filename}")81 dataset = load_dataset("lyi029/malvec-embeddings", split="train", streaming=True)82 embeddings = []83 count = 084 for sample in dataset:85 sample_filename = sample.get('source', '') or sample.get('filename', '')86 if filename.lower() in sample_filename.lower():87 emb = sample.get('embedding') or sample.get('emb')88 if emb is not None:89 embeddings.append(np.array(emb, dtype=np.float32))90 count += 191 if count >= 1000: break92 if embeddings:93 print(f"✅ Found {len(embeddings)} embeddings for {filename}")94 return embeddings95 else:96 print(f"⚠️ No embeddings found for {filename} in dataset")97 return None98 except Exception as e:99 print(f"❌ Error fetching from HF dataset: {e}")100 return None101 102def analyze_som_region_probability(103 som_analyzer: SOMAnalyzer,104 center_row: int,105 center_col: int,106 embeddings: List[np.ndarray],107 radius: int = 2108) -> Dict[str, float]:109 110 # (注意:你 main.py 裡的這個函數是同步的(def),不是(async def),這裡保持一致)111 print(f" 🔍 Analyzing region around ({center_row}, {center_col}) with radius={radius}")112 region_samples = []113 for emb in embeddings:114 sample_row, sample_col = som_analyzer.som.winner(emb) # 115 if abs(sample_row - center_row) <= radius and abs(sample_col - center_col) <= radius:116 region_samples.append(emb)117 print(f" Found {len(region_samples)} samples in region")118 if len(region_samples) == 0:119 return {"dropper": 0.0, "APT30": 0.0}120 121 feature_counts = {"dropper": [], "spreader": [], "APT30": [], "Codoso_Gh0st_1": []}122 if hasattr(som_analyzer, 'background_embeddings') and hasattr(som_analyzer, 'background_labels'):123 for i, bg_emb in enumerate(som_analyzer.background_embeddings):124 bg_row, bg_col = som_analyzer.som.winner(bg_emb) # 125 if abs(bg_row - center_row) <= radius and abs(bg_col - center_col) <= radius:126 for feature_name in feature_counts.keys():127 if feature_name in som_analyzer.background_labels:128 feature_counts[feature_name].append(129 som_analyzer.background_labels[feature_name][i]130 )131 probabilities = {}132 for feature_name, values in feature_counts.items():133 if len(values) > 0:134 prob = sum(values) / len(values)135 probabilities[feature_name] = float(prob)136 else:137 probabilities[feature_name] = 0.0138 return probabilities139 140 141async def analyze_with_som(embedding, predicted_family, filename):142 # ... (請從 main.py 完整複製貼上) ...143 if not SOM_AVAILABLE:144 print("⚠️ SOM analysis disabled")145 return None146 try:147 print("\n🗺️ Performing SOM analysis...")148 if isinstance(embedding, list):149 embedding_array = np.array(embedding, dtype=np.float32)150 else:151 embedding_array = embedding152 153 winner_row, winner_col = som_analyzer.find_winner_neuron(embedding_array)154 print(f" 📍 Winner neuron: ({winner_row}, {winner_col})")155 156 hf_embeddings = await get_embeddings_from_hf_dataset(filename)157 158 if hf_embeddings and len(hf_embeddings) > 0:159 print(f"✅ Using {len(hf_embeddings)} embeddings from HF dataset")160 features_prob = analyze_som_region_probability( # 161 som_analyzer, winner_row, winner_col, hf_embeddings, radius=2162 )163 data_source = "hf_dataset"164 else:165 print(f"⚠️ No HF data found, using background data")166 features_dict = som_analyzer.predict_features_from_position(167 winner_row, winner_col, radius=1168 )169 features_prob = {k: float(v) for k, v in features_dict.items()}170 data_source = "background_data"171 172 features_binary = {173 feature: 1 if prob > 0.2 else 0 174 for feature, prob in features_prob.items()175 }176 print(f" ✅ Features detected: {features_binary}")177 178 result = {179 "winner_position": {"row": int(winner_row), "col": int(col_winner)},180 "features": features_binary,181 "features_probability": features_prob,182 "data_source": data_source183 }184 print(f"✅ SOM analysis complete!")185 return result186 except Exception as e:187 print(f"❌ SOM analysis error: {e}")188 return None189 190@app.get("/")191def home():192 return {193 "message": "Malware Classifier API",194 "model": MODEL_REPO,195 "endpoints": {196 "/predict": "POST - Upload TXT files for prediction",197 "/model-info": "GET - Get model information"198 }199 }200 201@app.get("/model-info")202def get_model_info():203 """取得模型資訊"""204 # ✅ 檢查 id2label 的 key 類型205 id2label_sample = {}206 if hasattr(model.config, 'id2label'):207 # 取前 3 個 key 作為範例208 sample_keys = list(model.config.id2label.keys())[:3]209 id2label_sample = {210 k: {211 "type": str(type(k).__name__),212 "value": model.config.id2label[k]213 } for k in sample_keys214 }215 216 return {217 "model_name": MODEL_REPO,218 "num_labels": model.config.num_labels,219 "labels": list(model.config.id2label.values()) if hasattr(model.config, 'id2label') else None,220 "device": str(DEVICE),221 "max_length": 512,222 "id2label_sample": id2label_sample # ✅ 除錯用223 }224 225 226@app.post("/predict")227async def predict_malware(files: List[UploadFile] = File(...)):228 """229 批次預測惡意軟體家族230 231 接收多個 TXT 檔案 (來自同一個執行檔的分段)232 返回:233 - final_label: 多數投票的家族標籤234 - embeddings: 最高注意力權重的 embedding235 """236 237 if not files:238 raise HTTPException(status_code=400, detail="No files provided")239 240 print(f"\n{'='*60}")241 print(f"📥 Received {len(files)} files")242 print(f"{'='*60}")243 244 predictions = []245 all_embeddings = []246 247 try:248 with torch.no_grad():249 for idx, file in enumerate(files, 1):250 try:251 # 讀取檔案內容252 content = await file.read()253 asm_code = content.decode('utf-8').strip()254 255 if not asm_code:256 print(f"⚠️ Skip empty file: {file.filename}")257 continue258 259 print(f"📄 [{idx}/{len(files)}] Processing: {file.filename}")260 261 # Tokenize262 inputs = tokenizer(263 asm_code,264 return_tensors="pt",265 max_length=512,266 truncation=True,267 padding=True268 )269 inputs = {k: v.to(DEVICE) for k, v in inputs.items()}270 271 # 模型預測272 outputs = model(**inputs)273 274 # 取得預測結果275 logits = outputs.logits276 probs = torch.nn.functional.softmax(logits, dim=-1)277 confidence, predicted_class = torch.max(probs, dim=-1)278 279 # ✅ 取得 label (使用整數 key)280 predicted_id = predicted_class.item()281 if hasattr(model.config, 'id2label') and predicted_id in model.config.id2label:282 label_name = model.config.id2label[predicted_id]283 else:284 label_name = f"class_{predicted_id}"285 286 print(f" 🔍 Predicted ID: {predicted_id}, Label: {label_name}")287 288 # ✅ 取得 embedding (使用 hidden_states)289 if hasattr(outputs, 'hidden_states') and outputs.hidden_states is not None:290 # 取最後一層的 [CLS] token embedding291 last_hidden = outputs.hidden_states[-1] # (batch, seq_len, hidden_size)292 cls_embedding = last_hidden[:, 0, :].cpu().numpy()[0] # (768,)293 else:294 # 備用方案: 使用 logits295 cls_embedding = logits.cpu().numpy()[0]296 297 # ✅ 取得注意力權重298 max_attention_score = 0.0299 if hasattr(outputs, 'attentions') and outputs.attentions is not None:300 # 取最後一層的注意力,平均所有 head301 last_attention = outputs.attentions[-1] # (1, num_heads, seq_len, seq_len)302 avg_attention = last_attention.mean(dim=1) # (1, seq_len, seq_len)303 304 # 取 [CLS] token 對所有 token 的注意力305 cls_attention = avg_attention[0, 0, :].cpu().numpy() # (seq_len,)306 max_attention_score = float(cls_attention.max())307 308 # 儲存結果309 predictions.append({310 "filename": file.filename,311 "predicted_label": label_name,312 "confidence": float(confidence.item()),313 "attention_score": max_attention_score314 })315 316 all_embeddings.append({317 "filename": file.filename,318 "label": label_name,319 "embedding": cls_embedding.tolist(),320 "attention_score": max_attention_score321 })322 323 print(f" ✅ Predicted: {label_name} (confidence: {confidence.item():.3f})")324 325 except Exception as e:326 print(f" ❌ Error processing {file.filename}: {e}")327 continue328 329 if not predictions:330 raise HTTPException(status_code=400, detail="No valid files processed")331 332 # ===== 多數投票 (Majority Voting) =====333 label_votes = [p["predicted_label"] for p in predictions]334 vote_counts = Counter(label_votes)335 final_label = vote_counts.most_common(1)[0][0]336 final_count = vote_counts[final_label]337 338 print(f"\n📊 Vote distribution: {dict(vote_counts)}")339 print(f"🏆 Final prediction: {final_label} ({final_count}/{len(predictions)} votes)")340 341 # ===== 選擇最高注意力權重的 embedding (屬於最終預測的 label) =====342 # 過濾出預測為 final_label 的 embeddings343 same_label_embeddings = [344 emb for emb in all_embeddings 345 if emb["label"] == final_label346 ]347 348 if same_label_embeddings:349 # 選擇注意力權重最高的350 best_embedding = max(same_label_embeddings, key=lambda x: x["attention_score"])351 else:352 # 如果沒有,就選所有裡面注意力最高的353 best_embedding = max(all_embeddings, key=lambda x: x["attention_score"])354 355 print(f"🎯 Selected embedding from: {best_embedding['filename']}")356 print(f" Attention score: {best_embedding['attention_score']:.4f}")357 358 359 # ===== 返回結果 =====360 result = {361 "final_label": final_label,362 "vote_count": final_count,363 "total_segments": len(predictions),364 "confidence": sum(p["confidence"] for p in predictions) / len(predictions),365 "vote_distribution": dict(vote_counts),366 "embedding": {367 "values": best_embedding["embedding"],368 "source_file": best_embedding["filename"],369 "attention_score": best_embedding["attention_score"],370 "dimension": len(best_embedding["embedding"])371 },372 "segment_predictions": predictions373 }374 375 # ===== 新增:執行 SOM 分析 =====376 som_analysis = None377 final_embedding_values = result.get("embedding", {}).get("values")378 379 if SOM_AVAILABLE and final_embedding_values:380 print("\n🗺️ Performing SOM analysis on HF Space...")381 try:382 # 383 som_analysis = await analyze_with_som(384 embedding=final_embedding_values,385 predicted_family=result.get("final_label", "Unknown"),386 filename=files[0].filename # 387 )388 except Exception as e:389 print(f"❌ SOM analysis error on HF Space: {e}")390 391 # 392 result["som_analysis"] = som_analysis393 394 print(f"\n✅ Prediction with SOM complete!\n")395 396 return result397 398 except Exception as e:399 print(f"\n❌ Error in predict_malware: {e}")400 import traceback401 traceback.print_exc()402 raise HTTPException(status_code=500, detail=f"Prediction failed: {str(e)}")403 404@app.post("/predict-batch")405async def predict_batch(files: List[UploadFile] = File(...)):406 """407 簡化版: 只返回 label 和 embedding,不包含詳細資訊408 """409 result = await predict_malware(files)410 411 return {412 "final_label": result["final_label"],413 "embedding": result["embedding"]["values"]414 }415 416# ✅ HF Space 啟動入口417if __name__ == "__main__":418 import uvicorn419 port = int(os.environ.get("PORT", 7860))420 uvicorn.run(app, host="0.0.0.0", port=port)