Dina-Raslan commited on
Commit ·
eb9ae8c
1
Parent(s): 151e1d5
Add model and related files
Browse files- app.py +309 -0
- emotion_engine.py +52 -0
- emotion_model.pth +3 -0
- requirements.txt +0 -0
app.py
ADDED
|
@@ -0,0 +1,309 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from flask import Flask, request, jsonify, send_file
|
| 2 |
+
from flask_cors import CORS
|
| 3 |
+
import os
|
| 4 |
+
import uuid
|
| 5 |
+
import json
|
| 6 |
+
from datetime import datetime
|
| 7 |
+
import csv
|
| 8 |
+
import io
|
| 9 |
+
import numpy as np
|
| 10 |
+
import random
|
| 11 |
+
|
| 12 |
+
from config import Config
|
| 13 |
+
from models import db, Session, PageRecord, MouseEvent, FrameCapture
|
| 14 |
+
from ml_models.mouse_analyzer import MousePatternAnalyzer
|
| 15 |
+
from ml_models.fusion_model import RiskFusionModel
|
| 16 |
+
from ml_models.model_loader import model_loader
|
| 17 |
+
from ml_models.emotion_engine import EmotionEngine
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
app = Flask(__name__)
|
| 21 |
+
app.config.from_object(Config)
|
| 22 |
+
CORS(app)
|
| 23 |
+
|
| 24 |
+
db.init_app(app)
|
| 25 |
+
|
| 26 |
+
# تحميل النماذج المدربة
|
| 27 |
+
def load_trained_models():
|
| 28 |
+
models_loaded = model_loader.load_all_models(
|
| 29 |
+
emotion_path='ml_models/trained_models/emotion_model.pth',
|
| 30 |
+
mouse_path='ml_models/trained_models/mouse_model.pkl',
|
| 31 |
+
fusion_path='ml_models/trained_models/fusion_model.pkl'
|
| 32 |
+
)
|
| 33 |
+
|
| 34 |
+
if models_loaded:
|
| 35 |
+
print("🎉 All trained models loaded successfully!")
|
| 36 |
+
else:
|
| 37 |
+
print("⚠️ Some models failed to load, using fallback methods")
|
| 38 |
+
|
| 39 |
+
load_trained_models()
|
| 40 |
+
|
| 41 |
+
emotion_engine = EmotionEngine()
|
| 42 |
+
mouse_analyzer = MousePatternAnalyzer()
|
| 43 |
+
fusion_model = RiskFusionModel()
|
| 44 |
+
|
| 45 |
+
# صفحات الاختبار المحدثة (واقعية ومتنوعة)
|
| 46 |
+
TEST_PAGES = [
|
| 47 |
+
# صفحات البنوك
|
| 48 |
+
{'id': 'bank_secure', 'type': 'legitimate', 'url': '/pages/bank-secure', 'name': 'البنك الأهلي السعودي'},
|
| 49 |
+
{'id': 'bank_phishing', 'type': 'phishing', 'url': '/pages/bank-phishing', 'name': 'تنبيه البنك السعودي'},
|
| 50 |
+
|
| 51 |
+
# صفحات البريد الإلكتروني
|
| 52 |
+
{'id': 'email_secure', 'type': 'legitimate', 'url': '/pages/email-secure', 'name': 'Outlook - البريد'},
|
| 53 |
+
{'id': 'email_phishing', 'type': 'phishing', 'url': '/pages/email-phishing', 'name': 'رسالة خدمة العملاء'},
|
| 54 |
+
|
| 55 |
+
# صفحات وسائل التواصل
|
| 56 |
+
{'id': 'social_secure', 'type': 'legitimate', 'url': '/pages/social-secure', 'name': 'Facebook تسجيل الدخول'},
|
| 57 |
+
{'id': 'social_phishing', 'type': 'phishing', 'url': '/pages/social-phishing', 'name': 'تنبيه فيسبوك الأمني'},
|
| 58 |
+
|
| 59 |
+
# صفحات التسوق
|
| 60 |
+
{'id': 'shopping_secure', 'type': 'legitimate', 'url': '/pages/shopping-secure', 'name': 'Amazon تسجيل الدخول'},
|
| 61 |
+
{'id': 'shopping_phishing', 'type': 'phishing', 'url': '/pages/shopping-phishing', 'name': 'عرض خاص - خصم 80%'},
|
| 62 |
+
|
| 63 |
+
# صفحات الخدمات الحكومية
|
| 64 |
+
{'id': 'gov_secure', 'type': 'legitimate', 'url': '/pages/gov-secure', 'name': 'أبشر - الخدمات الإلكترونية'},
|
| 65 |
+
{'id': 'gov_phishing', 'type': 'phishing', 'url': '/pages/gov-phishing', 'name': 'تنبيه وزارة الداخلية'}
|
| 66 |
+
]
|
| 67 |
+
|
| 68 |
+
@app.route('/')
|
| 69 |
+
def home():
|
| 70 |
+
return jsonify({
|
| 71 |
+
'message': 'Phishing Study Backend API',
|
| 72 |
+
'status': 'running',
|
| 73 |
+
'endpoints': {
|
| 74 |
+
'health': '/api/health',
|
| 75 |
+
'start_session': '/api/session/start (POST)',
|
| 76 |
+
'submit_page': '/api/session/<session_id>/page (POST)',
|
| 77 |
+
'end_session': '/api/session/<session_id>/end (POST)',
|
| 78 |
+
'export_data': '/api/admin/export (GET)'
|
| 79 |
+
}
|
| 80 |
+
})
|
| 81 |
+
|
| 82 |
+
@app.route('/api/health', methods=['GET'])
|
| 83 |
+
def health_check():
|
| 84 |
+
return jsonify({'status': 'healthy', 'timestamp': datetime.utcnow().isoformat()})
|
| 85 |
+
|
| 86 |
+
@app.route('/api/session/start', methods=['POST'])
|
| 87 |
+
def start_session():
|
| 88 |
+
try:
|
| 89 |
+
data = request.get_json() or {}
|
| 90 |
+
user_id = data.get('user_id', f'user_{uuid.uuid4().hex[:8]}')
|
| 91 |
+
username = data.get('username')
|
| 92 |
+
session_id = f'sess_{uuid.uuid4().hex[:16]}'
|
| 93 |
+
pages_order = random.sample(TEST_PAGES, len(TEST_PAGES))
|
| 94 |
+
|
| 95 |
+
session = Session(
|
| 96 |
+
id=session_id,
|
| 97 |
+
user_id=user_id,
|
| 98 |
+
username=username,
|
| 99 |
+
consent_given=True,
|
| 100 |
+
pages_order=json.dumps(pages_order)
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
db.session.add(session)
|
| 104 |
+
db.session.commit()
|
| 105 |
+
|
| 106 |
+
return jsonify({
|
| 107 |
+
'session_id': session_id,
|
| 108 |
+
'user_id': user_id,
|
| 109 |
+
'username': username,
|
| 110 |
+
'pages': pages_order,
|
| 111 |
+
'start_time': session.start_time.isoformat()
|
| 112 |
+
}), 201
|
| 113 |
+
|
| 114 |
+
except Exception as e:
|
| 115 |
+
return jsonify({'error': str(e)}), 500
|
| 116 |
+
|
| 117 |
+
@app.route('/api/session/<session_id>/page', methods=['POST'])
|
| 118 |
+
def submit_page_data(session_id):
|
| 119 |
+
try:
|
| 120 |
+
data = request.get_json()
|
| 121 |
+
if not data:
|
| 122 |
+
return jsonify({'error': 'No data provided'}), 400
|
| 123 |
+
|
| 124 |
+
# ===== Session =====
|
| 125 |
+
session = Session.query.get(session_id)
|
| 126 |
+
if not session:
|
| 127 |
+
return jsonify({'error': 'Session not found'}), 404
|
| 128 |
+
|
| 129 |
+
# ===== Times =====
|
| 130 |
+
start_time = datetime.fromisoformat(data['start_time'].replace('Z', '+00:00'))
|
| 131 |
+
end_time = datetime.fromisoformat(data['end_time'].replace('Z', '+00:00'))
|
| 132 |
+
|
| 133 |
+
# ===== Frames =====
|
| 134 |
+
frames = data.get('frames', [])
|
| 135 |
+
if not frames:
|
| 136 |
+
return jsonify({'error': 'No frames received'}), 400
|
| 137 |
+
|
| 138 |
+
# ناخد آخر فريم بس (أنضف وأسرع)
|
| 139 |
+
last_frame = frames[-1]
|
| 140 |
+
base64_img = last_frame.get('img_base64')
|
| 141 |
+
|
| 142 |
+
if not base64_img:
|
| 143 |
+
return jsonify({'error': 'No image data in frame'}), 400
|
| 144 |
+
|
| 145 |
+
# ===== Emotion (PyTorch ViT Engine) =====
|
| 146 |
+
emotion_probs, dominant_emotion = emotion_engine.predict_from_base64(
|
| 147 |
+
base64_img
|
| 148 |
+
)
|
| 149 |
+
|
| 150 |
+
# حساب الـ emotion risk
|
| 151 |
+
emotion_risk = (
|
| 152 |
+
emotion_probs.get('fear', 0) * 0.6 +
|
| 153 |
+
emotion_probs.get('sad', 0) * 0.3 +
|
| 154 |
+
emotion_probs.get('angry', 0) * 0.1
|
| 155 |
+
)
|
| 156 |
+
emotion_risk = min(emotion_risk, 1.0)
|
| 157 |
+
|
| 158 |
+
# Debug (مفيد في التيرمنال)
|
| 159 |
+
print("======== EMOTION ENGINE DEBUG ========")
|
| 160 |
+
print("Dominant emotion:", dominant_emotion)
|
| 161 |
+
print("Emotion probabilities:")
|
| 162 |
+
for k, v in emotion_probs.items():
|
| 163 |
+
print(f" {k}: {v:.4f}")
|
| 164 |
+
print("Emotion risk:", emotion_risk)
|
| 165 |
+
print("=====================================")
|
| 166 |
+
|
| 167 |
+
# ===== Mouse =====
|
| 168 |
+
mouse_events = data.get('mouse_events', [])
|
| 169 |
+
mouse_features = mouse_analyzer.extract_features(
|
| 170 |
+
mouse_events, start_time, end_time
|
| 171 |
+
)
|
| 172 |
+
mouse_risk = mouse_analyzer.calculate_mouse_risk(mouse_features)
|
| 173 |
+
|
| 174 |
+
# ===== Fusion =====
|
| 175 |
+
phishing_score = fusion_model.calculate_phishing_score(
|
| 176 |
+
emotion_risk, mouse_risk, mouse_features
|
| 177 |
+
)
|
| 178 |
+
|
| 179 |
+
# ===== PageRecord =====
|
| 180 |
+
page_record = PageRecord(
|
| 181 |
+
id=f'page_rec_{uuid.uuid4().hex[:16]}',
|
| 182 |
+
session_id=session_id,
|
| 183 |
+
page_id=data.get('page_id'),
|
| 184 |
+
page_type=data.get('page_type'),
|
| 185 |
+
start_time=start_time,
|
| 186 |
+
end_time=end_time,
|
| 187 |
+
user_label=data.get('label'),
|
| 188 |
+
notes=data.get('notes', '')
|
| 189 |
+
)
|
| 190 |
+
|
| 191 |
+
page_record.emotion_probs = json.dumps(emotion_probs)
|
| 192 |
+
page_record.dominant_emotion = dominant_emotion
|
| 193 |
+
page_record.emotion_risk = emotion_risk
|
| 194 |
+
page_record.mouse_features = json.dumps(mouse_features)
|
| 195 |
+
page_record.mouse_risk = mouse_risk
|
| 196 |
+
page_record.phishing_score = phishing_score
|
| 197 |
+
|
| 198 |
+
db.session.add(page_record)
|
| 199 |
+
|
| 200 |
+
# ===== Mouse Events =====
|
| 201 |
+
for event in mouse_events:
|
| 202 |
+
mouse_event = MouseEvent(
|
| 203 |
+
page_record_id=page_record.id,
|
| 204 |
+
event_type=event['type'],
|
| 205 |
+
x=event.get('x'),
|
| 206 |
+
y=event.get('y'),
|
| 207 |
+
timestamp=datetime.fromisoformat(event['t'].replace('Z', '+00:00')),
|
| 208 |
+
additional_data=json.dumps({
|
| 209 |
+
k: v for k, v in event.items()
|
| 210 |
+
if k not in ['type', 'x', 'y', 't']
|
| 211 |
+
})
|
| 212 |
+
)
|
| 213 |
+
db.session.add(mouse_event)
|
| 214 |
+
|
| 215 |
+
db.session.commit()
|
| 216 |
+
|
| 217 |
+
# ===== Response =====
|
| 218 |
+
return jsonify({
|
| 219 |
+
'emotion_source': 'internal_engine',
|
| 220 |
+
'emotion_probs': emotion_probs,
|
| 221 |
+
'dominant_emotion': dominant_emotion,
|
| 222 |
+
'emotion_risk': emotion_risk,
|
| 223 |
+
'mouse_features': mouse_features,
|
| 224 |
+
'mouse_risk': mouse_risk,
|
| 225 |
+
'phishing_score': phishing_score,
|
| 226 |
+
'risk_level': fusion_model.get_risk_level(phishing_score)
|
| 227 |
+
}), 201
|
| 228 |
+
|
| 229 |
+
except Exception as e:
|
| 230 |
+
db.session.rollback()
|
| 231 |
+
return jsonify({'error': str(e)}), 500
|
| 232 |
+
|
| 233 |
+
|
| 234 |
+
@app.route('/api/admin/export', methods=['GET'])
|
| 235 |
+
def export_data():
|
| 236 |
+
try:
|
| 237 |
+
records = PageRecord.query.all()
|
| 238 |
+
|
| 239 |
+
output = io.StringIO()
|
| 240 |
+
writer = csv.writer(output)
|
| 241 |
+
|
| 242 |
+
writer.writerow([
|
| 243 |
+
'record_id', 'user_id', 'username', 'session_id', 'page_id', 'page_type',
|
| 244 |
+
'start_time', 'timestamp', 'sample_type', 'emotion_probs',
|
| 245 |
+
'mouse_features', 'emotion_risk', 'mouse_risk', 'phishing_score',
|
| 246 |
+
'label', 'dominant_emotion', 'notes'
|
| 247 |
+
])
|
| 248 |
+
|
| 249 |
+
for record in records:
|
| 250 |
+
writer.writerow([
|
| 251 |
+
record.id,
|
| 252 |
+
record.session.user_id,
|
| 253 |
+
record.session.username,
|
| 254 |
+
record.session_id,
|
| 255 |
+
record.page_id,
|
| 256 |
+
record.page_type,
|
| 257 |
+
record.start_time.isoformat() if record.start_time else '',
|
| 258 |
+
record.timestamp.isoformat() if record.timestamp else '',
|
| 259 |
+
'detail',
|
| 260 |
+
record.emotion_probs or '{}',
|
| 261 |
+
record.mouse_features or '{}',
|
| 262 |
+
record.emotion_risk or 0.0,
|
| 263 |
+
record.mouse_risk or 0.0,
|
| 264 |
+
record.phishing_score or 0.0,
|
| 265 |
+
record.user_label or '',
|
| 266 |
+
record.dominant_emotion or '',
|
| 267 |
+
record.notes or ''
|
| 268 |
+
])
|
| 269 |
+
|
| 270 |
+
output.seek(0)
|
| 271 |
+
return send_file(
|
| 272 |
+
io.BytesIO(output.getvalue().encode('utf-8-sig')),
|
| 273 |
+
mimetype='text/csv; charset=utf-8',
|
| 274 |
+
as_attachment=True,
|
| 275 |
+
download_name=f'phishing_study_export_{datetime.utcnow().strftime("%Y%m%d_%H%M%S")}.csv'
|
| 276 |
+
)
|
| 277 |
+
except Exception as e:
|
| 278 |
+
return jsonify({'error': str(e)}), 500
|
| 279 |
+
|
| 280 |
+
@app.route('/api/session/<session_id>', methods=['GET'])
|
| 281 |
+
def get_session(session_id):
|
| 282 |
+
session = Session.query.get(session_id)
|
| 283 |
+
if not session:
|
| 284 |
+
return jsonify({'error': 'Session not found'}), 404
|
| 285 |
+
|
| 286 |
+
page_records = PageRecord.query.filter_by(session_id=session_id).all()
|
| 287 |
+
|
| 288 |
+
return jsonify({
|
| 289 |
+
'session_id': session.id,
|
| 290 |
+
'user_id': session.user_id,
|
| 291 |
+
'username': session.username,
|
| 292 |
+
'start_time': session.start_time.isoformat(),
|
| 293 |
+
'end_time': session.end_time.isoformat() if session.end_time else None,
|
| 294 |
+
'completed': session.completed,
|
| 295 |
+
'page_records': [
|
| 296 |
+
{
|
| 297 |
+
'page_id': pr.page_id,
|
| 298 |
+
'user_label': pr.user_label,
|
| 299 |
+
'phishing_score': pr.phishing_score,
|
| 300 |
+
'dominant_emotion': pr.dominant_emotion
|
| 301 |
+
} for pr in page_records
|
| 302 |
+
]
|
| 303 |
+
})
|
| 304 |
+
|
| 305 |
+
if __name__ == '__main__':
|
| 306 |
+
with app.app_context():
|
| 307 |
+
db.create_all()
|
| 308 |
+
print("Phishing Study Backend starting on http://localhost:5000")
|
| 309 |
+
app.run(debug=True, host='0.0.0.0', port=5000)
|
emotion_engine.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import base64
|
| 3 |
+
import os
|
| 4 |
+
import io
|
| 5 |
+
from PIL import Image
|
| 6 |
+
from transformers import AutoImageProcessor, AutoConfig, AutoModelForImageClassification
|
| 7 |
+
|
| 8 |
+
class EmotionEngine:
|
| 9 |
+
def __init__(self):
|
| 10 |
+
# اسم الموديل مش مهم يظهر في أي حتة تانية
|
| 11 |
+
self.processor = AutoImageProcessor.from_pretrained(
|
| 12 |
+
"trpakov/vit-face-expression"
|
| 13 |
+
)
|
| 14 |
+
|
| 15 |
+
config = AutoConfig.from_pretrained(
|
| 16 |
+
"trpakov/vit-face-expression"
|
| 17 |
+
)
|
| 18 |
+
|
| 19 |
+
self.model = AutoModelForImageClassification.from_config(config)
|
| 20 |
+
|
| 21 |
+
# ... داخل الكلاس __init__
|
| 22 |
+
model_dir = os.path.join(os.path.dirname(__file__), "trained_models")
|
| 23 |
+
model_path = os.path.join(model_dir, "emotion_model.pth")
|
| 24 |
+
|
| 25 |
+
state_dict = torch.load(model_path, map_location="cpu")
|
| 26 |
+
self.model.load_state_dict(state_dict)
|
| 27 |
+
self.model.eval()
|
| 28 |
+
|
| 29 |
+
self.labels = self.model.config.id2label
|
| 30 |
+
|
| 31 |
+
def predict_from_base64(self, base64_img):
|
| 32 |
+
# فك الصورة
|
| 33 |
+
img_bytes = base64.b64decode(base64_img.split(",")[1])
|
| 34 |
+
img = Image.open(io.BytesIO(img_bytes)).convert("RGB")
|
| 35 |
+
|
| 36 |
+
# preprocessing
|
| 37 |
+
inputs = self.processor(images=img, return_tensors="pt")
|
| 38 |
+
|
| 39 |
+
with torch.no_grad():
|
| 40 |
+
outputs = self.model(**inputs)
|
| 41 |
+
probs = torch.softmax(outputs.logits, dim=1)[0]
|
| 42 |
+
|
| 43 |
+
emotion_probs = {
|
| 44 |
+
self.labels[i]: float(probs[i])
|
| 45 |
+
for i in range(len(probs))
|
| 46 |
+
}
|
| 47 |
+
|
| 48 |
+
dominant_emotion = max(
|
| 49 |
+
emotion_probs, key=emotion_probs.get
|
| 50 |
+
)
|
| 51 |
+
|
| 52 |
+
return emotion_probs, dominant_emotion
|
emotion_model.pth
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:c6b95ab1a233a3341db52aee0522934412f580449dea2978fdbc80fa70438336
|
| 3 |
+
size 343295867
|
requirements.txt
ADDED
|
Binary file (1.66 kB). View file
|
|
|