Java实现AB测试案例

wen java案例 2

本文目录导读:

Java实现AB测试案例

  1. 基础AB测试框架
  2. 统计分析功能
  3. 实验管理器
  4. 主程序示例
  5. Maven依赖配置

我将为您提供一个完整的AB测试实现案例,包含测试框架、数据分析和可视化。

基础AB测试框架

import java.util.*;
import java.util.concurrent.*;
import java.util.stream.*;
/**
 * AB测试核心框架
 */
public class ABTestFramework {
    // 测试配置
    public static class TestConfig {
        private String testName;
        private double trafficPercentage;  // 流量分配比例
        private int sampleSize;            // 最小样本量
        private double significanceLevel;  // 显著性水平(通常0.05)
        private double power;              // 统计功效(通常0.8)
        private double baselineConversion; // 基线转化率
        private double minimumDetectableEffect; // 最小可检测效果
        public TestConfig(String testName, double baselineConversion, double minimumDetectableEffect) {
            this.testName = testName;
            this.trafficPercentage = 1.0;  // 默认100%流量
            this.significanceLevel = 0.05;
            this.power = 0.8;
            this.baselineConversion = baselineConversion;
            this.minimumDetectableEffect = minimumDetectableEffect;
            this.sampleSize = calculateSampleSize();
        }
        private int calculateSampleSize() {
            // 使用标准公式计算样本量
            double zAlpha = 1.96; // 95%置信水平
            double zBeta = 0.84;  // 80%统计功效
            double p1 = baselineConversion;
            double p2 = baselineConversion + minimumDetectableEffect;
            double pBar = (p1 + p2) / 2;
            double sampleSize = Math.pow(zAlpha * Math.sqrt(2 * pBar * (1 - pBar)) 
                    + zBeta * Math.sqrt(p1 * (1 - p1) + p2 * (1 - p2)), 2) 
                    / Math.pow(p2 - p1, 2);
            return (int) Math.ceil(sampleSize);
        }
        // Getters
        public String getTestName() { return testName; }
        public int getSampleSize() { return sampleSize; }
        public double getSignificanceLevel() { return significanceLevel; }
    }
    // 用户分组
    public enum Group {
        CONTROL,    // 对照组
        TREATMENT   // 实验组
    }
    // 测试结果
    public static class TestResult {
        private int controlVisitors;
        private int controlConversions;
        private int treatmentVisitors;
        private int treatmentConversions;
        private double pValue;
        private boolean isSignificant;
        public double getControlConversionRate() {
            return controlVisitors == 0 ? 0 : (double) controlConversions / controlVisitors;
        }
        public double getTreatmentConversionRate() {
            return treatmentVisitors == 0 ? 0 : (double) treatmentConversions / treatmentVisitors;
        }
        public double getLift() {
            double controlRate = getControlConversionRate();
            return controlRate == 0 ? 0 : (getTreatmentConversionRate() - controlRate) / controlRate;
        }
        // Getters and Setters
        public int getControlVisitors() { return controlVisitors; }
        public void setControlVisitors(int controlVisitors) { this.controlVisitors = controlVisitors; }
        public int getControlConversions() { return controlConversions; }
        public void setControlConversions(int controlConversions) { this.controlConversions = controlConversions; }
        public int getTreatmentVisitors() { return treatmentVisitors; }
        public void setTreatmentVisitors(int treatmentVisitors) { this.treatmentVisitors = treatmentVisitors; }
        public int getTreatmentConversions() { return treatmentConversions; }
        public void setTreatmentConversions(int treatmentConversions) { this.treatmentConversions = treatmentConversions; }
        public double getPValue() { return pValue; }
        public void setPValue(double pValue) { this.pValue = pValue; }
        public boolean isSignificant() { return isSignificant; }
        public void setSignificant(boolean significant) { isSignificant = significant; }
    }
    // 用户分配器
    public static class UserAssigner {
        private final Random random;
        private final Map<String, Group> assignments;
        public UserAssigner() {
            this.random = new Random();
            this.assignments = new ConcurrentHashMap<>();
        }
        public Group assignUser(String userId, double treatmentProbability) {
            // 检查是否已分配
            if (assignments.containsKey(userId)) {
                return assignments.get(userId);
            }
            // 新用户分配
            Group group = random.nextDouble() < treatmentProbability ? 
                         Group.TREATMENT : Group.CONTROL;
            assignments.put(userId, group);
            return group;
        }
        public Group getUserGroup(String userId) {
            return assignments.getOrDefault(userId, Group.CONTROL);
        }
    }
}

统计分析功能

import org.apache.commons.math3.distribution.NormalDistribution;
import org.apache.commons.math3.distribution.ChiSquaredDistribution;
/**
 * 统计分析工具
 */
public class StatisticalAnalysis {
    // Z检验(双比例检验)
    public static class ZTestResult {
        public double zScore;
        public double pValue;
        public boolean isSignificant;
        @Override
        public String toString() {
            return String.format("Z-score: %.4f, p-value: %.4f, Significant: %b", 
                zScore, pValue, isSignificant);
        }
    }
    public static ZTestResult zTestForProportions(
            int treatmentSuccess, int treatmentTotal,
            int controlSuccess, int controlTotal,
            double alpha) {
        double treatmentRate = (double) treatmentSuccess / treatmentTotal;
        double controlRate = (double) controlSuccess / controlTotal;
        // 合并比例
        double pooledRate = (double) (treatmentSuccess + controlSuccess) / 
                          (treatmentTotal + controlTotal);
        // 标准误差
        double standardError = Math.sqrt(
            pooledRate * (1 - pooledRate) * (1.0/treatmentTotal + 1.0/controlTotal)
        );
        ZTestResult result = new ZTestResult();
        if (standardError == 0) {
            result.zScore = 0;
            result.pValue = 1;
            result.isSignificant = false;
            return result;
        }
        // 计算Z分数
        result.zScore = (treatmentRate - controlRate) / standardError;
        // 计算p值(双边检验)
        NormalDistribution normalDist = new NormalDistribution();
        result.pValue = 2 * (1 - normalDist.cumulativeProbability(Math.abs(result.zScore)));
        // 判断显著性
        result.isSignificant = result.pValue < alpha;
        return result;
    }
    // 卡方检验
    public static double chiSquareTest(int[][] contingencyTable) {
        int rows = contingencyTable.length;
        int cols = contingencyTable[0].length;
        // 计算行列总和
        int[] rowSums = new int[rows];
        int[] colSums = new int[cols];
        int total = 0;
        for (int i = 0; i < rows; i++) {
            for (int j = 0; j < cols; j++) {
                rowSums[i] += contingencyTable[i][j];
                colSums[j] += contingencyTable[i][j];
                total += contingencyTable[i][j];
            }
        }
        // 计算卡方统计量
        double chiSquare = 0;
        for (int i = 0; i < rows; i++) {
            for (int j = 0; j < cols; j++) {
                double expected = (double) (rowSums[i] * colSums[j]) / total;
                if (expected > 0) {
                    chiSquare += Math.pow(contingencyTable[i][j] - expected, 2) / expected;
                }
            }
        }
        // 计算p值
        int df = (rows - 1) * (cols - 1);
        ChiSquaredDistribution chiSqDist = new ChiSquaredDistribution(df);
        double pValue = 1 - chiSqDist.cumulativeProbability(chiSquare);
        return pValue;
    }
    // 计算置信区间
    public static double[] confidenceInterval(double success, double total, double confidenceLevel) {
        double rate = success / total;
        double z = (confidenceLevel == 0.95) ? 1.96 : 
                   (confidenceLevel == 0.99) ? 2.576 : 1.645;
        double margin = z * Math.sqrt(rate * (1 - rate) / total);
        return new double[] {rate - margin, rate + margin};
    }
    // 贝叶斯AB测试分析
    public static class BayesianAnalysis {
        public static double probabilityTreatmentBetter(
                int treatmentSuccess, int treatmentTotal,
                int controlSuccess, int controlTotal,
                int iterations) {
            int betterCount = 0;
            Random random = new Random();
            for (int i = 0; i < iterations; i++) {
                // 从Beta分布中抽样
                double treatmentSample = betaDistributionSample(
                    treatmentSuccess + 1, treatmentTotal - treatmentSuccess + 1, random);
                double controlSample = betaDistributionSample(
                    controlSuccess + 1, controlTotal - controlSuccess + 1, random);
                if (treatmentSample > controlSample) {
                    betterCount++;
                }
            }
            return (double) betterCount / iterations;
        }
        private static double betaDistributionSample(double alpha, double beta, Random random) {
            // 使用Bennett方法近似Beta分布抽样
            double u = random.nextDouble();
            double v = random.nextDouble();
            double x = Math.pow(u, 1.0/alpha);
            double y = Math.pow(v, 1.0/beta);
            while (x + y > 1) {
                u = random.nextDouble();
                v = random.nextDouble();
                x = Math.pow(u, 1.0/alpha);
                y = Math.pow(v, 1.0/beta);
            }
            return x / (x + y);
        }
    }
}

实验管理器

import java.time.*;
import java.util.*;
import java.util.concurrent.*;
import java.util.concurrent.atomic.*;
/**
 * AB测试管理器
 */
public class ABTestManager {
    private final Map<String, Experiment> activeExperiments;
    private final Map<String, TestResult> completedResults;
    private final UserAssigner userAssigner;
    public ABTestManager() {
        this.activeExperiments = new ConcurrentHashMap<>();
        this.completedResults = new ConcurrentHashMap<>();
        this.userAssigner = new UserAssigner();
    }
    // 实验配置
    public static class Experiment {
        private String id;
        private String name;
        private ABTestFramework.TestConfig config;
        private LocalDateTime startTime;
        private LocalDateTime endTime;
        private AtomicInteger controlVisitors = new AtomicInteger();
        private AtomicInteger controlConversions = new AtomicInteger();
        private AtomicInteger treatmentVisitors = new AtomicInteger();
        private AtomicInteger treatmentConversions = new AtomicInteger();
        public Experiment(String id, String name, ABTestFramework.TestConfig config) {
            this.id = id;
            this.name = name;
            this.config = config;
            this.startTime = LocalDateTime.now();
        }
        // Getters and helper methods
        public void incrementVisitor(ABTestFramework.Group group) {
            if (group == ABTestFramework.Group.CONTROL) {
                controlVisitors.incrementAndGet();
            } else {
                treatmentVisitors.incrementAndGet();
            }
        }
        public void incrementConversion(ABTestFramework.Group group) {
            if (group == ABTestFramework.Group.CONTROL) {
                controlConversions.incrementAndGet();
            } else {
                treatmentConversions.incrementAndGet();
            }
        }
        // Getters...
        public String getId() { return id; }
        public String getName() { return name; }
        public ABTestFramework.TestConfig getConfig() { return config; }
        public int getControlVisitors() { return controlVisitors.get(); }
        public int getControlConversions() { return controlConversions.get(); }
        public int getTreatmentVisitors() { return treatmentVisitors.get(); }
        public int getTreatmentConversions() { return treatmentConversions.get(); }
    }
    // 创建实验
    public Experiment createExperiment(String id, String name, 
            double baselineConversion, double minDetectableEffect) {
        ABTestFramework.TestConfig config = new ABTestFramework.TestConfig(
            name, baselineConversion, minDetectableEffect);
        Experiment experiment = new Experiment(id, name, config);
        activeExperiments.put(id, experiment);
        return experiment;
    }
    // 记录访问
    public ABTestFramework.Group recordVisit(String experimentId, String userId) {
        Experiment experiment = activeExperiments.get(experimentId);
        if (experiment == null) {
            throw new IllegalArgumentException("Experiment not found: " + experimentId);
        }
        ABTestFramework.Group group = userAssigner.assignUser(userId, 0.5);
        experiment.incrementVisitor(group);
        return group;
    }
    // 记录转化
    public void recordConversion(String experimentId, String userId) {
        Experiment experiment = activeExperiments.get(experimentId);
        if (experiment == null) {
            throw new IllegalArgumentException("Experiment not found: " + experimentId);
        }
        ABTestFramework.Group group = userAssigner.getUserGroup(userId);
        experiment.incrementConversion(group);
    }
    // 完成实验并获取结果
    public ABTestFramework.TestResult completeExperiment(String experimentId) {
        Experiment experiment = activeExperiments.get(experimentId);
        if (experiment == null) {
            throw new IllegalArgumentException("Experiment not found: " + experimentId);
        }
        // 构建结果
        ABTestFramework.TestResult result = new ABTestFramework.TestResult();
        result.setControlVisitors(experiment.getControlVisitors());
        result.setControlConversions(experiment.getControlConversions());
        result.setTreatmentVisitors(experiment.getTreatmentVisitors());
        result.setTreatmentConversions(experiment.getTreatmentConversions());
        // 进行统计分析
        StatisticalAnalysis.ZTestResult zTest = StatisticalAnalysis.zTestForProportions(
            experiment.getTreatmentConversions(),
            experiment.getTreatmentVisitors(),
            experiment.getControlConversions(),
            experiment.getControlVisitors(),
            experiment.getConfig().getSignificanceLevel()
        );
        result.setPValue(zTest.pValue);
        result.setSignificant(zTest.isSignificant);
        // 更新实验状态
        experiment.endTime = LocalDateTime.now();
        completedResults.put(experimentId, result);
        activeExperiments.remove(experimentId);
        return result;
    }
    // 获取实验报告
    public String generateReport(String experimentId) {
        Experiment experiment = null;
        ABTestFramework.TestResult result = completedResults.get(experimentId);
        if (result == null) {
            // 尝试从active experiments获取
            experiment = activeExperiments.get(experimentId);
            if (experiment == null) {
                return "Experiment not found";
            }
            // 生成实时报告
            StringBuilder sb = new StringBuilder();
            sb.append("=== 实验进行中报告 ===\n");
            sb.append("实验名: ").append(experiment.getName()).append("\n");
            sb.append("时长: ").append(Duration.between(experiment.startTime, 
                LocalDateTime.now()).toHours()).append(" 小时\n");
            sb.append("对照组: ").append(experiment.getControlVisitors())
              .append(" 访客, ").append(experiment.getControlConversions())
              .append(" 转化\n");
            double controlRate = experiment.getControlVisitors() > 0 ? 
                (double) experiment.getControlConversions() / experiment.getControlVisitors() : 0;
            sb.append(String.format("对照组转化率: %.4f%%\n", controlRate * 100));
            sb.append("实验组: ").append(experiment.getTreatmentVisitors())
              .append(" 访客, ").append(experiment.getTreatmentConversions())
              .append(" 转化\n");
            double treatmentRate = experiment.getTreatmentVisitors() > 0 ? 
                (double) experiment.getTreatmentConversions() / experiment.getTreatmentVisitors() : 0;
            sb.append(String.format("实验组转化率: %.4f%%\n", treatmentRate * 100));
            return sb.toString();
        }
        // 生成最终报告
        return generateFinalReport(experiment, result);
    }
    private String generateFinalReport(Experiment experiment, 
            ABTestFramework.TestResult result) {
        StringBuilder sb = new StringBuilder();
        sb.append("=== 最终AB测试报告 ===\n");
        sb.append("实验名: ").append(experiment.getName()).append("\n");
        sb.append("开始时间: ").append(experiment.startTime).append("\n");
        sb.append("结束时间: ").append(experiment.endTime).append("\n\n");
        sb.append("对照组: ").append(result.getControlVisitors())
          .append(" 访客, ").append(result.getControlConversions())
          .append(" 转化\n");
        sb.append(String.format("对照组转化率: %.4f%%\n", 
            result.getControlConversionRate() * 100));
        sb.append("实验组: ").append(result.getTreatmentVisitors())
          .append(" 访客, ").append(result.getTreatmentConversions())
          .append(" 转化\n");
        sb.append(String.format("实验组转化率: %.4f%%\n", 
            result.getTreatmentConversionRate() * 100));
        sb.append(String.format("提升幅度: %.2f%%\n", result.getLift() * 100));
        sb.append(String.format("P值: %.4f\n", result.getPValue()));
        sb.append(" ").append(result.isSignificant() ? 
            "统计显著" : "统计不显著").append("\n");
        // 添加额外分析
        double[] controlCI = StatisticalAnalysis.confidenceInterval(
            result.getControlConversions(), result.getControlVisitors(), 0.95);
        double[] treatmentCI = StatisticalAnalysis.confidenceInterval(
            result.getTreatmentConversions(), result.getTreatmentVisitors(), 0.95);
        sb.append(String.format("对照组95%%置信区间: [%.4f, %.4f]\n", 
            controlCI[0] * 100, controlCI[1] * 100));
        sb.append(String.format("实验组95%%置信区间: [%.4f, %.4f]\n", 
            treatmentCI[0] * 100, treatmentCI[1] * 100));
        return sb.toString();
    }
}

主程序示例

import java.util.*;
import java.util.concurrent.*;
/**
 * 主程序示例
 */
public class ABTestMain {
    public static void main(String[] args) throws InterruptedException {
        // 创建AB测试管理器
        ABTestManager manager = new ABTestManager();
        // 创建实验
        // 假设:当前转化率为5%,希望检测2%的提升
        ABTestManager.Experiment experiment = manager.createExperiment(
            "TEST-001", 
            "新首页设计测试",
            0.05,    // 基线转化率5%
            0.02     // 最小可检测效果2%
        );
        System.out.println("实验创建成功!");
        System.out.println("信息: " + experiment.getName());
        System.out.println("ID: " + experiment.getId());
        System.out.println("需要样本量: " + experiment.getConfig().getSampleSize());
        System.out.println();
        // 模拟用户流量
        Random random = new Random();
        int totalUsers = 10000;
        int conversionRate = 5; // 基准转化率5%
        int treatmentBoost = 2; // 实验组额外提升2%
        System.out.println("开始模拟流量...");
        for (int i = 1; i <= totalUsers; i++) {
            String userId = "user_" + i;
            // 记录访问
            ABTestFramework.Group group = manager.recordVisit("TEST-001", userId);
            // 模拟转化
            int rate = group == ABTestFramework.Group.TREATMENT ? 
                      conversionRate + treatmentBoost : conversionRate;
            if (random.nextInt(100) < rate) {
                manager.recordConversion("TEST-001", userId);
            }
            // 显示进度
            if (i % 1000 == 0) {
                System.out.printf("已处理 %d/%d 用户%n", i, totalUsers);
            }
        }
        // 完成实验
        System.out.println("\n完成实验...");
        ABTestFramework.TestResult result = manager.completeExperiment("TEST-001");
        // 打印最终报告
        System.out.println(manager.generateReport("TEST-001"));
        // 额外统计
        System.out.println("\n=== 额外分析 ===");
        // 贝叶斯分析
        double probability = StatisticalAnalysis.BayesianAnalysis
            .probabilityTreatmentBetter(
                result.getTreatmentConversions(), 
                result.getTreatmentVisitors(),
                result.getControlConversions(),
                result.getControlVisitors(),
                10000);
        System.out.printf("贝叶斯概率(实验组优于对照组): %.2f%%\n", 
            probability * 100);
        // 卡方检验
        int[][] contingencyTable = {
            {result.getTreatmentConversions(), 
             result.getTreatmentVisitors() - result.getTreatmentConversions()},
            {result.getControlConversions(), 
             result.getControlVisitors() - result.getControlConversions()}
        };
        double chiSquarePValue = StatisticalAnalysis.chiSquareTest(contingencyTable);
        System.out.printf("卡方检验p值: %.4f\n", chiSquarePValue);
        // 样本量建议
        System.out.printf("\n建议: 最小样本量 %d 人%n", 
            experiment.getConfig().getSampleSize());
    }
}

Maven依赖配置

<dependencies>
    <!-- Apache Commons Math for statistical functions -->
    <dependency>
        <groupId>org.apache.commons</groupId>
        <artifactId>commons-math3</artifactId>
        <version>3.6.1</version>
    </dependency>
    <!-- MySQL connector for storing results -->
    <dependency>
        <groupId>mysql</groupId>
        <artifactId>mysql-connector-java</artifactId>
        <version>8.0.33</version>
    </dependency>
    <!-- Logging -->
    <dependency>
        <groupId>org.slf4j</groupId>
        <artifactId>slf4j-api</artifactId>
        <version>2.0.7</version>
    </dependency>
</dependencies>

这个完整的AB测试框架包含:

  1. 核心框架:实验配置、用户分组、结果记录
  2. 统计分析:Z检验、卡方检验、置信区间、贝叶斯分析
  3. 实验管理:创建实验、分配用户、生成报告
  4. 实时监控:实验进行中的数据跟踪

可以根据实际需求扩展此框架,例如添加:

  • 多变量测试支持
  • 基于机器学习的分层抽样
  • 实时监控仪表板
  • 自动停止规则(当达到显著性时)

抱歉,评论功能暂时关闭!