Java协同过滤推荐案例

wen java案例 3

本文目录导读:

Java协同过滤推荐案例

  1. 项目结构
  2. 完整代码实现
  3. 使用说明

我来提供一个完整的Java协同过滤推荐系统案例,包含基于用户的协同过滤和基于物品的协同过滤两种算法。

项目结构

src/main/java/com/recommend/
├── RecommendationSystem.java  // 主类
├── model/
│   ├── UserRating.java       // 用户评分模型
│   └── Recommendation.java   // 推荐结果模型
├── algorithm/
│   ├── UserBasedCF.java      // 基于用户的协同过滤
│   └── ItemBasedCF.java      // 基于物品的协同过滤
└── util/
    └── SimilarityUtil.java   // 相似度计算工具

完整代码实现

数据模型类

package com.recommend.model;
import java.util.HashMap;
import java.util.Map;
/**
 * 用户评分模型
 */
public class UserRating {
    // 用户ID -> 物品ID -> 评分
    private Map<String, Map<String, Double>> userRatings;
    public UserRating() {
        this.userRatings = new HashMap<>();
    }
    /**
     * 添加用户评分
     * @param userId 用户ID
     * @param itemId 物品ID
     * @param rating 评分(1-5分)
     */
    public void addRating(String userId, String itemId, double rating) {
        userRatings.computeIfAbsent(userId, k -> new HashMap<>())
                   .put(itemId, rating);
    }
    /**
     * 获取用户的所有评分
     */
    public Map<String, Double> getUserRatings(String userId) {
        return userRatings.getOrDefault(userId, new HashMap<>());
    }
    /**
     * 获取所有用户
     */
    public Map<String, Map<String, Double>> getAllUserRatings() {
        return userRatings;
    }
    /**
     * 获取用户评分
     */
    public double getRating(String userId, String itemId) {
        Map<String, Double> ratings = userRatings.get(userId);
        if (ratings != null && ratings.containsKey(itemId)) {
            return ratings.get(itemId);
        }
        return 0; // 表示没评分
    }
    /**
     * 获取所有评价过某个物品的用户
     */
    public Map<String, Double> getItemRatings(String itemId) {
        Map<String, Double> itemUsers = new HashMap<>();
        for (Map.Entry<String, Map<String, Double>> entry : userRatings.entrySet()) {
            String userId = entry.getKey();
            Map<String, Double> ratings = entry.getValue();
            if (ratings.containsKey(itemId)) {
                itemUsers.put(userId, ratings.get(itemId));
            }
        }
        return itemUsers;
    }
    /**
     * 获取所有物品
     */
    public java.util.Set<String> getAllItems() {
        java.util.Set<String> items = new java.util.HashSet<>();
        for (Map<String, Double> ratings : userRatings.values()) {
            items.addAll(ratings.keySet());
        }
        return items;
    }
}

推荐结果类

package com.recommend.model;
/**
 * 推荐结果
 */
public class Recommendation implements Comparable<Recommendation> {
    private String itemId;
    private double score;
    public Recommendation(String itemId, double score) {
        this.itemId = itemId;
        this.score = score;
    }
    public String getItemId() {
        return itemId;
    }
    public double getScore() {
        return score;
    }
    @Override
    public int compareTo(Recommendation o) {
        // 按得分降序排列
        return Double.compare(o.score, this.score);
    }
    @Override
    public String toString() {
        return String.format("Recommendation{item='%s', score=%.4f}", itemId, score);
    }
}

相似度计算工具类

package com.recommend.util;
import java.util.Map;
import java.util.Set;
/**
 * 相似度计算工具
 */
public class SimilarityUtil {
    /**
     * 计算Pearson相关系数
     * @param user1Ratings 用户1的评分
     * @param user2Ratings 用户2的评分
     */
    public static double pearsonCorrelation(Map<String, Double> user1Ratings, 
                                          Map<String, Double> user2Ratings) {
        // 找到共同评分的物品
        Set<String> commonItems = new java.util.HashSet<>(user1Ratings.keySet());
        commonItems.retainAll(user2Ratings.keySet());
        if (commonItems.size() < 2) {
            return 0.0; // 共同评分物品太少,无法计算
        }
        double sum1 = 0, sum2 = 0, sum1Sq = 0, sum2Sq = 0, pSum = 0;
        int n = commonItems.size();
        for (String item : commonItems) {
            double r1 = user1Ratings.get(item);
            double r2 = user2Ratings.get(item);
            sum1 += r1;
            sum2 += r2;
            sum1Sq += r1 * r1;
            sum2Sq += r2 * r2;
            pSum += r1 * r2;
        }
        double num = pSum - (sum1 * sum2 / n);
        double den = Math.sqrt((sum1Sq - sum1 * sum1 / n) * (sum2Sq - sum2 * sum2 / n));
        if (den == 0) return 0.0;
        return num / den;
    }
    /**
     * 计算余弦相似度
     */
    public static double cosineSimilarity(Map<String, Double> user1Ratings, 
                                        Map<String, Double> user2Ratings) {
        Set<String> commonItems = new java.util.HashSet<>(user1Ratings.keySet());
        commonItems.retainAll(user2Ratings.keySet());
        if (commonItems.isEmpty()) {
            return 0.0;
        }
        double dotProduct = 0;
        double norm1 = 0;
        double norm2 = 0;
        // 计算共同物品的点积
        for (String item : commonItems) {
            dotProduct += user1Ratings.get(item) * user2Ratings.get(item);
        }
        // 计算各自的范数
        for (double rating : user1Ratings.values()) {
            norm1 += rating * rating;
        }
        for (double rating : user2Ratings.values()) {
            norm2 += rating * rating;
        }
        if (norm1 == 0 || norm2 == 0) {
            return 0.0;
        }
        return dotProduct / (Math.sqrt(norm1) * Math.sqrt(norm2));
    }
    /**
     * 修正的余弦相似度(用于物品间相似度)
     */
    public static double adjustedCosineSimilarity(Map<String, Double> item1Ratings, 
                                                Map<String, Double> item2Ratings,
                                                Map<String, Double> userAvgRatings) {
        Set<String> commonUsers = new java.util.HashSet<>(item1Ratings.keySet());
        commonUsers.retainAll(item2Ratings.keySet());
        if (commonUsers.isEmpty()) {
            return 0.0;
        }
        double sum = 0;
        double sum1Sq = 0;
        double sum2Sq = 0;
        for (String user : commonUsers) {
            double avg = userAvgRatings.getOrDefault(user, 3.0); // 默认平均分3.0
            double r1 = item1Ratings.get(user) - avg;
            double r2 = item2Ratings.get(user) - avg;
            sum += r1 * r2;
            sum1Sq += r1 * r1;
            sum2Sq += r2 * r2;
        }
        if (sum1Sq == 0 || sum2Sq == 0) {
            return 0.0;
        }
        return sum / (Math.sqrt(sum1Sq) * Math.sqrt(sum2Sq));
    }
}

基于用户的协同过滤算法

package com.recommend.algorithm;
import com.recommend.model.Recommendation;
import com.recommend.model.UserRating;
import com.recommend.util.SimilarityUtil;
import java.util.*;
import java.util.stream.Collectors;
/**
 * 基于用户的协同过滤算法
 */
public class UserBasedCF {
    private UserRating userRating;
    private Map<String, Map<String, Double>> userSimilarityCache;
    public UserBasedCF(UserRating userRating) {
        this.userRating = userRating;
        this.userSimilarityCache = new HashMap<>();
    }
    /**
     * 为用户生成推荐
     * @param userId 用户ID
     * @param topN 推荐数量
     * @param k 邻居数量
     */
    public List<Recommendation> recommend(String userId, int topN, int k) {
        Map<String, Double> targetUserRatings = userRating.getUserRatings(userId);
        Map<String, Map<String, Double>> allUsers = userRating.getAllUserRatings();
        // 1. 计算与其他用户的相似度
        Map<String, Double> userSimilarities = new HashMap<>();
        for (String otherUser : allUsers.keySet()) {
            if (otherUser.equals(userId)) continue;
            double similarity = SimilarityUtil.pearsonCorrelation(
                targetUserRatings, 
                allUsers.get(otherUser)
            );
            if (similarity > 0) { // 只考虑正相似度
                userSimilarities.put(otherUser, similarity);
            }
        }
        // 2. 找到K个最近邻居
        List<Map.Entry<String, Double>> nearestNeighbors = userSimilarities.entrySet()
            .stream()
            .sorted(Map.Entry.<String, Double>comparingByValue().reversed())
            .limit(k)
            .collect(Collectors.toList());
        // 3. 生成推荐
        Map<String, Double> itemScores = new HashMap<>();
        Map<String, Double> itemScoreWeight = new HashMap<>();
        for (Map.Entry<String, Double> neighbor : nearestNeighbors) {
            String neighborUser = neighbor.getKey();
            double similarity = neighbor.getValue();
            Map<String, Double> neighborRatings = userRating.getUserRatings(neighborUser);
            // 计算邻居用户的平均评分
            double neighborAvg = neighborRatings.values().stream()
                .mapToDouble(Double::doubleValue)
                .average()
                .orElse(3.0);
            for (Map.Entry<String, Double> rating : neighborRatings.entrySet()) {
                String itemId = rating.getKey();
                double itemRating = rating.getValue();
                // 只推荐用户没有评分的物品
                if (!targetUserRatings.containsKey(itemId)) {
                    // 基于相似度和评分的加权平均
                    double weightedScore = similarity * (itemRating - neighborAvg);
                    itemScores.put(itemId, 
                        itemScores.getOrDefault(itemId, 0.0) + weightedScore);
                    itemScoreWeight.put(itemId, 
                        itemScoreWeight.getOrDefault(itemId, 0.0) + Math.abs(similarity));
                }
            }
        }
        // 4. 计算最终推荐分数
        List<Recommendation> recommendations = new ArrayList<>();
        for (Map.Entry<String, Double> entry : itemScores.entrySet()) {
            String itemId = entry.getKey();
            double weight = itemScoreWeight.getOrDefault(itemId, 0.0);
            if (weight > 0) {
                // 加上用户平均分作为基准
                double userAvg = targetUserRatings.values().stream()
                    .mapToDouble(Double::doubleValue)
                    .average()
                    .orElse(3.0);
                double score = userAvg + (entry.getValue() / weight);
                recommendations.add(new Recommendation(itemId, score));
            }
        }
        // 按得分排序,取前N个
        return recommendations.stream()
            .sorted(Comparator.reverseOrder())
            .limit(topN)
            .collect(Collectors.toList());
    }
    /**
     * 获取与指定用户最相似的用户
     */
    public List<String> getSimilarUsers(String userId, int k) {
        Map<String, Double> targetUserRatings = userRating.getUserRatings(userId);
        Map<String, Double> similarities = new HashMap<>();
        for (Map.Entry<String, Map<String, Double>> entry : 
                userRating.getAllUserRatings().entrySet()) {
            String otherUser = entry.getKey();
            if (!otherUser.equals(userId)) {
                double sim = SimilarityUtil.pearsonCorrelation(
                    targetUserRatings, entry.getValue());
                similarities.put(otherUser, sim);
            }
        }
        return similarities.entrySet().stream()
            .sorted(Map.Entry.<String, Double>comparingByValue().reversed())
            .limit(k)
            .map(Map.Entry::getKey)
            .collect(Collectors.toList());
    }
}

基于物品的协同过滤算法

package com.recommend.algorithm;
import com.recommend.model.Recommendation;
import com.recommend.model.UserRating;
import com.recommend.util.SimilarityUtil;
import java.util.*;
import java.util.stream.Collectors;
/**
 * 基于物品的协同过滤算法
 */
public class ItemBasedCF {
    private UserRating userRating;
    private Map<String, Map<String, Double>> itemSimilarityCache;
    public ItemBasedCF(UserRating userRating) {
        this.userRating = userRating;
        this.itemSimilarityCache = new HashMap<>();
    }
    /**
     * 为用户生成推荐
     * @param userId 目标用户
     * @param topN 推荐数量
     * @param k 相似物品数量
     */
    public List<Recommendation> recommend(String userId, int topN, int k) {
        Map<String, Double> userRatings = userRating.getUserRatings(userId);
        Set<String> allItems = userRating.getAllItems();
        // 计算用户平均评分
        double userAvg = userRatings.values().stream()
            .mapToDouble(Double::doubleValue)
            .average()
            .orElse(3.0);
        // 对每个未评分的物品,预测评分
        Map<String, Double> predictions = new HashMap<>();
        for (String item : allItems) {
            if (!userRatings.containsKey(item)) {
                double predictedScore = predictRating(userId, item, k);
                if (predictedScore > 0) {
                    predictions.put(item, predictedScore);
                }
            }
        }
        // 获取当前用户的评分物品列表
        List<Map.Entry<String, Double>> ratedItems = new ArrayList<>(userRatings.entrySet());
        // 生成推荐
        List<Recommendation> recommendations = new ArrayList<>();
        for (Map.Entry<String, Double> prediction : predictions.entrySet()) {
            String targetItem = prediction.getKey();
            double score = prediction.getValue();
            // 可选:可以结合用户历史评分调整
            recommendations.add(new Recommendation(targetItem, score));
        }
        // 按得分排序,取前N个
        return recommendations.stream()
            .sorted(Comparator.reverseOrder())
            .limit(topN)
            .collect(Collectors.toList());
    }
    /**
     * 预测用户对物品的评分
     */
    private double predictRating(String userId, String itemId, int k) {
        Map<String, Double> userRatings = userRating.getUserRatings(userId);
        // 找到与目标物品最相似的k个物品
        List<Map.Entry<String, Double>> similarItems = getSimilarItems(itemId, k);
        if (similarItems.isEmpty()) {
            return 0.0;
        }
        double sum = 0.0;
        double sumWeight = 0.0;
        for (Map.Entry<String, Double> similarItem : similarItems) {
            String similarItemId = similarItem.getKey();
            double similarity = similarItem.getValue();
            double itemRating = userRatings.getOrDefault(similarItemId, 0.0);
            if (itemRating > 0) {
                sum += similarity * itemRating;
                sumWeight += similarity;
            }
        }
        if (sumWeight == 0) {
            return 0.0;
        }
        return sum / sumWeight;
    }
    /**
     * 获取与指定物品最相似的物品列表
     */
    public List<Map.Entry<String, Double>> getSimilarItems(String itemId, int k) {
        if (itemSimilarityCache.containsKey(itemId)) {
            Map<String, Double> cachedSims = itemSimilarityCache.get(itemId);
            return cachedSims.entrySet().stream()
                .sorted(Map.Entry.<String, Double>comparingByValue().reversed())
                .limit(k)
                .collect(Collectors.toList());
        }
        Map<String, Double> itemRatings = userRating.getItemRatings(itemId);
        Set<String> allItems = userRating.getAllItems();
        // 计算用户平均评分(用于修正余弦相似度)
        Map<String, Double> userAvgRatings = new HashMap<>();
        for (Map.Entry<String, Map<String, Double>> entry : 
                userRating.getAllUserRatings().entrySet()) {
            double avg = entry.getValue().values().stream()
                .mapToDouble(Double::doubleValue)
                .average()
                .orElse(3.0);
            userAvgRatings.put(entry.getKey(), avg);
        }
        Map<String, Double> similarities = new HashMap<>();
        for (String otherItem : allItems) {
            if (!otherItem.equals(itemId)) {
                Map<String, Double> otherItemRatings = userRating.getItemRatings(otherItem);
                double sim = SimilarityUtil.adjustedCosineSimilarity(
                    itemRatings, otherItemRatings, userAvgRatings);
                if (sim > 0) {
                    similarities.put(otherItem, sim);
                }
            }
        }
        // 缓存相似度
        itemSimilarityCache.put(itemId, similarities);
        // 返回top-k相似物品
        return similarities.entrySet().stream()
            .sorted(Map.Entry.<String, Double>comparingByValue().reversed())
            .limit(k)
            .collect(Collectors.toList());
    }
    /**
     * 获取与指定物品最相似的物品
     */
    public List<String> getMostSimilarItems(String itemId, int k) {
        return getSimilarItems(itemId, k).stream()
            .map(Map.Entry::getKey)
            .collect(Collectors.toList());
    }
}

推荐系统主类

package com.recommend;
import com.recommend.algorithm.ItemBasedCF;
import com.recommend.algorithm.UserBasedCF;
import com.recommend.model.Recommendation;
import com.recommend.model.UserRating;
import java.util.List;
import java.util.Scanner;
/**
 * 推荐系统主类
 */
public class RecommendationSystem {
    private UserRating userRating;
    private UserBasedCF userBasedCF;
    private ItemBasedCF itemBasedCF;
    /**
     * 初始化推荐系统
     */
    public void init() {
        // 加载模拟数据
        loadTestData();
        // 初始化算法
        userBasedCF = new UserBasedCF(userRating);
        itemBasedCF = new ItemBasedCF(userRating);
    }
    /**
     * 加载测试数据
     */
    private void loadTestData() {
        userRating = new UserRating();
        // 模拟用户评分数据 (5个用户, 6个物品)
        // 用户1:喜欢书籍和音乐(评分1-5)
        userRating.addRating("user1", "book1", 5.0);
        userRating.addRating("user1", "book2", 4.0);
        userRating.addRating("user1", "music1", 5.0);
        userRating.addRating("user1", "movie1", 3.0);
        userRating.addRating("user1", "sport1", 1.0);
        userRating.addRating("user1", "food1", 2.0);
        // 用户2:喜欢音乐
        userRating.addRating("user2", "music1", 5.0);
        userRating.addRating("user2", "music2", 4.0);
        userRating.addRating("user2", "book1", 4.0);
        userRating.addRating("user2", "movie2", 2.0);
        userRating.addRating("user2", "food1", 3.0);
        // 用户3:喜欢体育
        userRating.addRating("user3", "sport1", 5.0);
        userRating.addRating("user3", "sport2", 4.0);
        userRating.addRating("user3", "book2", 2.0);
        userRating.addRating("user3", "movie1", 1.0);
        // 用户4:喜欢电影和书籍
        userRating.addRating("user4", "movie1", 5.0);
        userRating.addRating("user4", "movie2", 4.0);
        userRating.addRating("user4", "book2", 4.0);
        userRating.addRating("user4", "book1", 3.0);
        userRating.addRating("user4", "music2", 2.0);
        // 用户5:综合爱好者
        userRating.addRating("user5", "book1", 4.0);
        userRating.addRating("user5", "music1", 3.0);
        userRating.addRating("user5", "sport2", 3.0);
        userRating.addRating("user5", "food1", 4.0);
        userRating.addRating("user5", "movie1", 3.0);
    }
    /**
     * 获取数据统计信息
     */
    public void printStatistics() {
        System.out.println("========== 数据集统计 ==========");
        System.out.println("用户数量: " + userRating.getAllUserRatings().size());
        System.out.println("物品数量: " + userRating.getAllItems().size());
        // 打印所有评分
        System.out.println("\n评分矩阵:");
        System.out.println("------------------------------------------------------------------------");
        for (Map.Entry<String, Map<String, Double>> user : 
                userRating.getAllUserRatings().entrySet()) {
            System.out.printf("%-10s", user.getKey() + ": ");
            for (Map.Entry<String, Double> rating : user.getValue().entrySet()) {
                System.out.printf("%s=%.1f ", rating.getKey(), rating.getValue());
            }
            System.out.println();
        }
        System.out.println("------------------------------------------------------------------------\n");
    }
    /**
     * 运行基于用户的协同过滤推荐
     * @param userId 目标用户
     */
    public void runUserBasedCF(String userId, int topN, int k) {
        System.out.println("\n基于用户的协同过滤推荐 (User-Based Collaborative Filtering)");
        System.out.println("========================================");
        List<Recommendation> recommendations = userBasedCF.recommend(userId, topN, k);
        if (recommendations.isEmpty()) {
            System.out.println("没有找到合适的推荐");
            return;
        }
        System.out.printf("用户 %s 的 Top-%d 推荐:", userId, topN);
        System.out.println("\n推荐排名:");
        for (int i = 0; i < recommendations.size(); i++) {
            Recommendation rec = recommendations.get(i);
            System.out.printf("%d. 物品 %s, 预测评分: %.2f\n", 
                i+1, rec.getItemId(), rec.getScore());
        }
    }
    /**
     * 运行基于物品的协同过滤推荐
     * @param userId 目标用户
     */
    public void runItemBasedCF(String userId, int topN, int k) {
        System.out.println("\n基于物品的协同过滤推荐 (Item-Based Collaborative Filtering)");
        System.out.println("============================================");
        List<Recommendation> recommendations = itemBasedCF.recommend(userId, topN, k);
        if (recommendations.isEmpty()) {
            System.out.println("没有找到合适的推荐");
            return;
        }
        System.out.printf("用户 %s 的 Top-%d 推荐:", userId, topN);
        System.out.println("\n推荐排名:");
        for (int i = 0; i < recommendations.size(); i++) {
            Recommendation rec = recommendations.get(i);
            System.out.printf("%d. 物品 %s, 预测评分: %.2f\n", 
                i+1, rec.getItemId(), rec.getScore());
        }
    }
    /**
     * 显示相似用户
     */
    public void showSimilarUsers(String userId, int k) {
        System.out.println("\n相似用户 (" + userId + "):");
        System.out.println("==============");
        List<String> similarUsers = userBasedCF.getSimilarUsers(userId, k);
        for (String similarUser : similarUsers) {
            System.out.println("- " + similarUser);
        }
    }
    /**
     * 显示相似物品
     */
    public void showSimilarItems(String itemId, int k) {
        System.out.println("\n相似物品 (" + itemId + "):");
        System.out.println("=============");
        List<String> similarItems = itemBasedCF.getMostSimilarItems(itemId, k);
        for (String similarItem : similarItems) {
            System.out.println("- " + similarItem);
        }
    }
    /**
     * 简单交互界面
     */
    public void interactiveMode() {
        Scanner scanner = new Scanner(System.in);
        while (true) {
            System.out.println("\n========== 协同过滤推荐系统 ==========");
            System.out.println("1. 显示数据集统计");
            System.out.println("2. 基于用户的协同过滤推荐");
            System.out.println("3. 基于物品的协同过滤推荐");
            System.out.println("4. 查看相似用户");
            System.out.println("5. 查看相似物品");
            System.out.println("6. 退出");
            System.out.print("请选择操作: ");
            int choice = scanner.nextInt();
            scanner.nextLine();
            switch (choice) {
                case 1:
                    printStatistics();
                    break;
                case 2:
                    System.out.print("请输入用户ID (user1-user5): ");
                    String userId = scanner.nextLine();
                    runUserBasedCF(userId, 5, 3);
                    break;
                case 3:
                    System.out.print("请输入用户ID (user1-user5): ");
                    userId = scanner.nextLine();
                    runItemBasedCF(userId, 5, 3);
                    break;
                case 4:
                    System.out.print("请输入用户ID (user1-user5): ");
                    userId = scanner.nextLine();
                    showSimilarUsers(userId, 2);
                    break;
                case 5:
                    System.out.print("请输入物品ID: ");
                    String itemId = scanner.nextLine();
                    showSimilarItems(itemId, 3);
                    break;
                case 6:
                    System.out.println("感谢使用,再见!");
                    return;
                default:
                    System.out.println("无效选择,请重试!");
            }
        }
    }
    /**
     * 主入口函数
     */
    public static void main(String[] args) {
        RecommendationSystem system = new RecommendationSystem();
        // 1. 初始化推荐系统
        system.init();
        // 2. 显示数据统计
        system.printStatistics();
        // 3. 运行推荐示例
        String testUser = "user1";
        // 用户协同过滤
        system.runUserBasedCF(testUser, 5, 3);
        // 物品协同过滤
        system.runItemBasedCF(testUser, 5, 3);
        // 4. 显示相似用户和物品
        system.showSimilarUsers(testUser, 3);
        system.showSimilarItems("book1", 3);
        // 5. 进入交互模式
        system.interactiveMode();
    }
}

Maven配置文件 (pom.xml)

<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
         xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
         xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 
         http://maven.apache.org/xsd/maven-4.0.0.xsd">
    <modelVersion>4.0.0</modelVersion>
    <groupId>com.example</groupId>
    <artifactId>recommendation-system</artifactId>
    <version>1.0-SNAPSHOT</version>
    <packaging>jar</packaging>
    <properties>
        <maven.compiler.source>8</maven.compiler.source>
        <maven.compiler.target>8</maven.compiler.target>
        <project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
    </properties>
    <dependencies>
        <!-- 如果需要日志 -->
        <dependency>
            <groupId>org.slf4j</groupId>
            <artifactId>slf4j-api</artifactId>
            <version>1.7.30</version>
        </dependency>
        <dependency>
            <groupId>org.slf4j</groupId>
            <artifactId>slf4j-simple</artifactId>
            <version>1.7.30</version>
        </dependency>
        <!-- 单元测试 -->
        <dependency>
            <groupId>junit</groupId>
            <artifactId>junit</artifactId>
            <version>4.13.2</version>
            <scope>test</scope>
        </dependency>
    </dependencies>
</project>

测试代码

package com.recommend;
import com.recommend.model.Recommendation;
import org.junit.Before;
import org.junit.Test;
import java.util.List;
import static org.junit.Assert.*;
public class RecommendationSystemTest {
    private RecommendationSystem system;
    @Before
    public void setUp() {
        system = new RecommendationSystem();
        system.init();
    }
    @Test
    public void testUserBasedRecommendation() {
        List<Recommendation> recommendations = system.userBasedCF.recommend("user1", 5, 3);
        assertNotNull("推荐结果不应为空", recommendations);
        assertFalse("推荐列表不应为空", recommendations.isEmpty());
        assertTrue("推荐数量不应超过5", recommendations.size() <= 5);
        // 验证推荐分数
        for (Recommendation rec : recommendations) {
            assertTrue("推荐分数应在1-5之间", rec.getScore() >= 1 && rec.getScore() <= 5);
        }
    }
    @Test
    public void testSimilarUserFound() {
        List<String> similarUsers = system.userBasedCF.getSimilarUsers("user1", 3);
        assertFalse("应找到相似用户", similarUsers.isEmpty());
        assertFalse("不应该包含自己", similarUsers.contains("user1"));
    }
    @Test
    public void testItemBasedRecommendation() {
        List<Recommendation> recommendations = system.itemBasedCF.recommend("user1", 5, 3);
        assertNotNull("推荐结果不应为空", recommendations);
        assertFalse("推荐列表不应为空", recommendations.isEmpty());
    }
}

使用说明

  1. 编译运行

    mvn clean compile
    mvn exec:java -Dexec.mainClass="com.recommend.RecommendationSystem"
  2. 功能特点

    • 基于用户的协同过滤(发现相似用户)
    • 基于物品的协同过滤(发现相似物品)
    • 两种算法结果对比
    • 交互式命令行界面
  3. 数据说明

    • 5个用户对6个物品的评分
    • 评分为1-5的整数
    • 测试数据模拟了不同的用户兴趣

这个案例实现了完整的协同过滤推荐系统,包括数据模型、相似度计算、推荐算法和交互界面,可以直接运行测试。

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