Team Ai
Apppublic

sameersyed/Defence_FrameWork

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
dataset_analyzer.py387 linesDownload Raw Back to backend
1import pandas as pd2import numpy as np3import matplotlib4matplotlib.use('Agg')5import matplotlib.pyplot as plt6import seaborn as sns7from io import BytesIO8import base649import json10 11class DatasetAnalyzer:12    def __init__(self):13        self.risk_thresholds = {14            'poisoning': 0.7,15            'imbalance': 0.6,16            'outliers': 0.5,17            'duplicates': 0.418        }19    20    def analyze_dataset(self, df):21        """Comprehensive dataset security analysis"""22        results = {23            'basic_info': self._get_basic_info(df),24            'risks': self._detect_risks(df),25            'class_distribution': self._analyze_classes(df),26            'outliers': self._detect_outliers(df),27            'duplicates': self._check_duplicates(df),28            'missing_values': self._analyze_missing(df),29            'statistics': self._get_statistics(df)30        }31        32        # Calculate overall risk score33        results['overall_risk'] = self._calculate_risk_score(results['risks'])34        35        return results36    37    def _get_basic_info(self, df):38        return {39            'total_samples': len(df),40            'total_features': len(df.columns),41            'numeric_features': len(df.select_dtypes(include=[np.number]).columns),42            'categorical_features': len(df.select_dtypes(include=['object']).columns),43            'memory_usage': f"{df.memory_usage(deep=True).sum() / 1024:.2f} KB"44        }45    46    def _detect_risks(self, df):47        risks = []48        49        # Check for label poisoning indicators50        if len(df.columns) > 0:51            last_col = df.iloc[:, -1]52            if last_col.dtype in ['object', 'int64']:53                class_counts = last_col.value_counts()54                if len(class_counts) > 1:55                    ratio = class_counts.min() / class_counts.max()56                    if ratio < 0.1:57                        risks.append({58                            'type': 'Class Imbalance',59                            'severity': 'HIGH',60                            'score': 85,61                            'description': f'Severe class imbalance detected (ratio: {ratio:.3f})'62                        })63        64        # Check for duplicate samples (potential poisoning)65        dup_ratio = df.duplicated().sum() / len(df)66        if dup_ratio > 0.05:67            risks.append({68                'type': 'Duplicate Samples',69                'severity': 'MEDIUM',70                'score': 60,71                'description': f'{dup_ratio*100:.1f}% duplicate samples found'72            })73        74        # Check for missing values75        missing_ratio = df.isnull().sum().sum() / (len(df) * len(df.columns))76        if missing_ratio > 0.1:77            risks.append({78                'type': 'Missing Data',79                'severity': 'MEDIUM',80                'score': 55,81                'description': f'{missing_ratio*100:.1f}% missing values'82            })83        84        # Check for outliers in numeric columns85        numeric_cols = df.select_dtypes(include=[np.number]).columns86        outlier_count = 087        for col in numeric_cols:88            Q1 = df[col].quantile(0.25)89            Q3 = df[col].quantile(0.75)90            IQR = Q3 - Q191            outliers = ((df[col] < (Q1 - 1.5 * IQR)) | (df[col] > (Q3 + 1.5 * IQR))).sum()92            outlier_count += outliers93        94        outlier_ratio = outlier_count / (len(df) * len(numeric_cols)) if len(numeric_cols) > 0 else 095        if outlier_ratio > 0.05:96            risks.append({97                'type': 'Outliers Detected',98                'severity': 'MEDIUM',99                'score': 50,100                'description': f'{outlier_ratio*100:.1f}% outlier values detected'101            })102        103        if not risks:104            risks.append({105                'type': 'No Major Risks',106                'severity': 'LOW',107                'score': 10,108                'description': 'Dataset appears clean'109            })110        111        return risks112    113    def _analyze_classes(self, df):114        if len(df.columns) == 0:115            return {}116        117        last_col = df.iloc[:, -1]118        if last_col.dtype in ['object', 'int64']:119            return last_col.value_counts().to_dict()120        return {}121    122    def _detect_outliers(self, df):123        numeric_cols = df.select_dtypes(include=[np.number]).columns124        outliers = {}125        126        for col in numeric_cols[:5]:  # Limit to first 5 numeric columns127            Q1 = df[col].quantile(0.25)128            Q3 = df[col].quantile(0.75)129            IQR = Q3 - Q1130            outlier_mask = (df[col] < (Q1 - 1.5 * IQR)) | (df[col] > (Q3 + 1.5 * IQR))131            outliers[col] = int(outlier_mask.sum())132        133        return outliers134    135    def _check_duplicates(self, df):136        return {137            'total_duplicates': int(df.duplicated().sum()),138            'percentage': round(df.duplicated().sum() / len(df) * 100, 2)139        }140    141    def _analyze_missing(self, df):142        missing = df.isnull().sum()143        return {col: int(count) for col, count in missing.items() if count > 0}144    145    def _get_statistics(self, df):146        numeric_df = df.select_dtypes(include=[np.number])147        if len(numeric_df.columns) == 0:148            return {}149        150        stats = numeric_df.describe().to_dict()151        return {k: {stat: round(v, 2) for stat, v in vals.items()} for k, vals in stats.items()}152    153    def _calculate_risk_score(self, risks):154        if not risks:155            return 0156        157        total_score = sum(risk['score'] for risk in risks)158        return min(total_score, 100)159    160    def generate_charts(self, df, analysis):161        """Generate visualization charts"""162        charts = {}163        164        # 1. Risk Distribution Pie Chart165        charts['risk_pie'] = self._create_risk_pie(analysis['risks'])166        167        # 2. Class Distribution Bar Chart168        if analysis['class_distribution']:169            charts['class_bar'] = self._create_class_bar(analysis['class_distribution'])170        171        # 3. Missing Values Heatmap172        if analysis['missing_values']:173            charts['missing_heatmap'] = self._create_missing_heatmap(df)174        175        # 4. Outlier Detection Box Plot176        charts['outlier_box'] = self._create_outlier_box(df)177        178        # 5. Risk Score Gauge179        charts['risk_gauge'] = self._create_risk_gauge(analysis['overall_risk'])180        181        # 6. Feature Correlation Heatmap182        numeric_df = df.select_dtypes(include=[np.number])183        if len(numeric_df.columns) > 1:184            charts['correlation'] = self._create_correlation_heatmap(numeric_df)185        186        return charts187    188    def _create_risk_pie(self, risks):189        fig, ax = plt.subplots(figsize=(8, 6))190        labels = [r['type'] for r in risks]191        sizes = [r['score'] for r in risks]192        colors = ['#ff6b6b' if r['severity'] == 'HIGH' else '#ffa500' if r['severity'] == 'MEDIUM' else '#51cf66' for r in risks]193        194        ax.pie(sizes, labels=labels, colors=colors, autopct='%1.1f%%', startangle=90)195        ax.set_title('Risk Distribution', fontsize=14, fontweight='bold')196        197        return self._fig_to_base64(fig)198    199    def _create_class_bar(self, class_dist):200        fig, ax = plt.subplots(figsize=(10, 6))201        classes = list(class_dist.keys())[:10]  # Limit to 10 classes202        counts = [class_dist[c] for c in classes]203        204        bars = ax.bar(range(len(classes)), counts, color='#667eea')205        ax.set_xlabel('Class', fontsize=12)206        ax.set_ylabel('Sample Count', fontsize=12)207        ax.set_title('Class Distribution', fontsize=14, fontweight='bold')208        ax.set_xticks(range(len(classes)))209        ax.set_xticklabels(classes, rotation=45, ha='right')210        211        # Add value labels on bars212        for bar in bars:213            height = bar.get_height()214            ax.text(bar.get_x() + bar.get_width()/2., height,215                   f'{int(height)}', ha='center', va='bottom')216        217        plt.tight_layout()218        return self._fig_to_base64(fig)219    220    def _create_missing_heatmap(self, df):221        fig, ax = plt.subplots(figsize=(10, 6))222        missing = df.isnull().sum()223        missing = missing[missing > 0][:10]  # Top 10 columns with missing values224        225        if len(missing) > 0:226            ax.barh(range(len(missing)), missing.values, color='#ff6b6b')227            ax.set_yticks(range(len(missing)))228            ax.set_yticklabels(missing.index)229            ax.set_xlabel('Missing Count', fontsize=12)230            ax.set_title('Missing Values by Feature', fontsize=14, fontweight='bold')231            plt.tight_layout()232        233        return self._fig_to_base64(fig)234    235    def _create_outlier_box(self, df):236        numeric_cols = df.select_dtypes(include=[np.number]).columns[:5]237        238        if len(numeric_cols) > 0:239            fig, ax = plt.subplots(figsize=(10, 6))240            df[numeric_cols].boxplot(ax=ax)241            ax.set_title('Outlier Detection (Box Plot)', fontsize=14, fontweight='bold')242            ax.set_ylabel('Value', fontsize=12)243            plt.xticks(rotation=45, ha='right')244            plt.tight_layout()245            return self._fig_to_base64(fig)246        247        return None248    249    def _create_risk_gauge(self, risk_score):250        fig, ax = plt.subplots(figsize=(8, 4), subplot_kw={'projection': 'polar'})251        252        theta = np.linspace(0, np.pi, 100)253        r = np.ones(100)254        255        # Color gradient based on risk256        if risk_score < 30:257            color = '#51cf66'258            label = 'LOW RISK'259        elif risk_score < 60:260            color = '#ffa500'261            label = 'MEDIUM RISK'262        else:263            color = '#ff6b6b'264            label = 'HIGH RISK'265        266        ax.plot(theta, r, color='lightgray', linewidth=20)267        ax.plot(theta[:int(risk_score)], r[:int(risk_score)], color=color, linewidth=20)268        ax.set_ylim(0, 1.5)269        ax.set_yticks([])270        ax.set_xticks([])271        ax.text(np.pi/2, 0.5, f'{risk_score}%\n{label}', ha='center', va='center', 272               fontsize=16, fontweight='bold')273        ax.set_title('Overall Risk Score', fontsize=14, fontweight='bold', pad=20)274        275        return self._fig_to_base64(fig)276    277    def _create_correlation_heatmap(self, numeric_df):278        fig, ax = plt.subplots(figsize=(10, 8))279        corr = numeric_df.corr()280        281        sns.heatmap(corr, annot=True, fmt='.2f', cmap='coolwarm', center=0,282                   square=True, linewidths=1, cbar_kws={"shrink": 0.8}, ax=ax)283        ax.set_title('Feature Correlation Matrix', fontsize=14, fontweight='bold')284        plt.tight_layout()285        286        return self._fig_to_base64(fig)287    288    def _fig_to_base64(self, fig):289        buf = BytesIO()290        fig.savefig(buf, format='png', dpi=100, bbox_inches='tight')291        buf.seek(0)292        img_base64 = base64.b64encode(buf.read()).decode('utf-8')293        plt.close(fig)294        return f"data:image/png;base64,{img_base64}"295    296    def generate_report_html(self, analysis, charts):297        """Generate downloadable HTML report"""298        html = f"""299<!DOCTYPE html>300<html>301<head>302    <title>Dataset Security Analysis Report</title>303    <style>304        body {{ font-family: Arial, sans-serif; margin: 40px; background: #f5f5f5; }}305        .container {{ max-width: 1200px; margin: 0 auto; background: white; padding: 30px; border-radius: 10px; }}306        h1 {{ color: #667eea; border-bottom: 3px solid #667eea; padding-bottom: 10px; }}307        h2 {{ color: #333; margin-top: 30px; }}308        .risk-high {{ color: #ff6b6b; font-weight: bold; }}309        .risk-medium {{ color: #ffa500; font-weight: bold; }}310        .risk-low {{ color: #51cf66; font-weight: bold; }}311        .info-box {{ background: #f0f0f0; padding: 15px; border-radius: 5px; margin: 10px 0; }}312        .chart {{ margin: 20px 0; text-align: center; }}313        .chart img {{ max-width: 100%; border: 1px solid #ddd; border-radius: 5px; }}314        table {{ width: 100%; border-collapse: collapse; margin: 20px 0; }}315        th, td {{ padding: 12px; text-align: left; border-bottom: 1px solid #ddd; }}316        th {{ background: #667eea; color: white; }}317    </style>318</head>319<body>320    <div class="container">321        <h1>๐Ÿ”’ Dataset Security Analysis Report</h1>322        <p><strong>Generated:</strong> {pd.Timestamp.now().strftime('%Y-%m-%d %H:%M:%S')}</p>323        324        <h2>๐Ÿ“Š Basic Information</h2>325        <div class="info-box">326            <p><strong>Total Samples:</strong> {analysis['basic_info']['total_samples']}</p>327            <p><strong>Total Features:</strong> {analysis['basic_info']['total_features']}</p>328            <p><strong>Numeric Features:</strong> {analysis['basic_info']['numeric_features']}</p>329            <p><strong>Categorical Features:</strong> {analysis['basic_info']['categorical_features']}</p>330            <p><strong>Memory Usage:</strong> {analysis['basic_info']['memory_usage']}</p>331        </div>332        333        <h2>โš ๏ธ Risk Assessment</h2>334        <div class="info-box">335            <p><strong>Overall Risk Score:</strong> <span class="risk-{'high' if analysis['overall_risk'] > 60 else 'medium' if analysis['overall_risk'] > 30 else 'low'}">{analysis['overall_risk']}%</span></p>336        </div>337        338        <h3>Detected Risks:</h3>339        <table>340            <tr><th>Risk Type</th><th>Severity</th><th>Score</th><th>Description</th></tr>341"""342        343        for risk in analysis['risks']:344            html += f"""345            <tr>346                <td>{risk['type']}</td>347                <td class="risk-{risk['severity'].lower()}">{risk['severity']}</td>348                <td>{risk['score']}</td>349                <td>{risk['description']}</td>350            </tr>351"""352        353        html += """354        </table>355        356        <h2>๐Ÿ“ˆ Visualizations</h2>357"""358        359        for chart_name, chart_data in charts.items():360            if chart_data:361                html += f"""362        <div class="chart">363            <h3>{chart_name.replace('_', ' ').title()}</h3>364            <img src="{chart_data}" alt="{chart_name}">365        </div>366"""367        368        html += """369        <h2>๐Ÿ›ก๏ธ Recommendations</h2>370        <ul>371            <li>Remove duplicate samples to prevent data poisoning</li>372            <li>Balance class distribution using SMOTE or undersampling</li>373            <li>Handle missing values with imputation or removal</li>374            <li>Detect and remove outliers using IQR method</li>375            <li>Validate data sources and apply sanitization</li>376            <li>Use adversarial training for robust models</li>377        </ul>378        379        <p style="margin-top: 40px; text-align: center; color: #888;">380            Generated by ShieldML - Deep Learning Security Analysis System381        </p>382    </div>383</body>384</html>385"""386        return html387