Team Ai
Apppublic

malvec/codebert_Malvec

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
app.py420 linesDownload Raw Back to root
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)