Java实现推荐系统案例

wen java案例 2

本文目录导读:

Java实现推荐系统案例

  1. 项目结构
  2. 数据模型
  3. 协同过滤算法
  4. 基于内容的推荐
  5. 混合推荐系统
  6. 数据加载器
  7. 主程序
  8. 示例数据文件
  9. 高级特性扩展
  10. 使用方法
  11. 特点

我来为您提供一个完整的Java推荐系统实现案例,包括协同过滤算法、内容推荐和混合推荐。

项目结构

recommendation-system/
├── src/main/java/com/recommend/
│   ├── model/
│   │   ├── User.java
│   │   ├── Item.java
│   │   └── Rating.java
│   ├── data/
│   │   └── DataLoader.java
│   ├── algorithm/
│   │   ├── CollaborativeFiltering.java
│   │   ├── ContentBasedFiltering.java
│   │   └── HybridRecommender.java
│   └── Main.java
└── data/
    ├── users.csv
    ├── items.csv
    └── ratings.csv

数据模型

// User.java
package com.recommend.model;
import java.util.HashMap;
import java.util.Map;
public class User {
    private int id;
    private String name;
    private Map<String, Double> preferences; // 用户偏好特征
    private Map<Integer, Double> ratings;    // 用户评分记录
    public User(int id, String name) {
        this.id = id;
        this.name = name;
        this.preferences = new HashMap<>();
        this.ratings = new HashMap<>();
    }
    // Getters and Setters
    public int getId() { return id; }
    public void setId(int id) { this.id = id; }
    public String getName() { return name; }
    public void setName(String name) { this.name = name; }
    public Map<String, Double> getPreferences() { return preferences; }
    public void setPreferences(Map<String, Double> preferences) { this.preferences = preferences; }
    public Map<Integer, Double> getRatings() { return ratings; }
    public void setRatings(Map<Integer, Double> ratings) { this.ratings = ratings; }
    public void addRating(int itemId, double rating) {
        ratings.put(itemId, rating);
    }
    public void addPreference(String feature, double value) {
        preferences.put(feature, value);
    }
}
// Item.java
package com.recommend.model;
import java.util.HashMap;
import java.util.Map;
public class Item {
    private int id;
    private String name;
    private String category;
    private Map<String, Double> features; // 项目特征向量
    public Item(int id, String name, String category) {
        this.id = id;
        this.name = name;
        this.category = category;
        this.features = new HashMap<>();
    }
    // Getters and Setters
    public int getId() { return id; }
    public void setId(int id) { this.id = id; }
    public String getName() { return name; }
    public void setName(String name) { this.name = name; }
    public String getCategory() { return category; }
    public void setCategory(String category) { this.category = category; }
    public Map<String, Double> getFeatures() { return features; }
    public void setFeatures(Map<String, Double> features) { this.features = features; }
    public void addFeature(String feature, double value) {
        features.put(feature, value);
    }
}
// Rating.java
package com.recommend.model;
public class Rating {
    private int userId;
    private int itemId;
    private double score;
    public Rating(int userId, int itemId, double score) {
        this.userId = userId;
        this.itemId = itemId;
        this.score = score;
    }
    // Getters and Setters
    public int getUserId() { return userId; }
    public void setUserId(int userId) { this.userId = userId; }
    public int getItemId() { return itemId; }
    public void setItemId(int itemId) { this.itemId = itemId; }
    public double getScore() { return score; }
    public void setScore(double score) { this.score = score; }
}

协同过滤算法

// CollaborativeFiltering.java
package com.recommend.algorithm;
import com.recommend.model.User;
import com.recommend.model.Item;
import java.util.*;
public class CollaborativeFiltering {
    private Map<Integer, User> users;
    private Map<Integer, Item> items;
    public CollaborativeFiltering(Map<Integer, User> users, Map<Integer, Item> items) {
        this.users = users;
        this.items = items;
    }
    /**
     * 基于用户的协同过滤推荐
     */
    public Map<Integer, Double> recommendByUser(int userId, int topN) {
        User targetUser = users.get(userId);
        if (targetUser == null) return Collections.emptyMap();
        // 1. 计算目标用户与其他用户的相似度
        Map<Integer, Double> userSimilarities = new HashMap<>();
        for (User otherUser : users.values()) {
            if (otherUser.getId() != userId) {
                double similarity = calculateUserSimilarity(targetUser, otherUser);
                userSimilarities.put(otherUser.getId(), similarity);
            }
        }
        // 2. 找出最相似的K个用户
        List<Map.Entry<Integer, Double>> sortedSimilarities = 
            new ArrayList<>(userSimilarities.entrySet());
        sortedSimilarities.sort((a, b) -> Double.compare(b.getValue(), a.getValue()));
        int k = Math.min(10, sortedSimilarities.size());
        Set<Integer> similarUsers = new HashSet<>();
        for (int i = 0; i < k; i++) {
            similarUsers.add(sortedSimilarities.get(i).getKey());
        }
        // 3. 预测目标用户对未评分项目的评分
        Map<Integer, Double> predictions = new HashMap<>();
        for (Item item : items.values()) {
            if (targetUser.getRatings().containsKey(item.getId())) continue;
            double totalSimilarity = 0;
            double weightedRating = 0;
            for (int similarUserId : similarUsers) {
                User similarUser = users.get(similarUserId);
                Double rating = similarUser.getRatings().get(item.getId());
                if (rating != null) {
                    double similarity = userSimilarities.get(similarUserId);
                    totalSimilarity += similarity;
                    weightedRating += similarity * rating;
                }
            }
            if (totalSimilarity > 0) {
                predictions.put(item.getId(), weightedRating / totalSimilarity);
            }
        }
        // 4. 返回Top N推荐
        return getTopN(predictions, topN);
    }
    /**
     * 基于项目的协同过滤推荐
     */
    public Map<Integer, Double> recommendByItem(int userId, int topN) {
        User targetUser = users.get(userId);
        if (targetUser == null) return Collections.emptyMap();
        Map<Integer, Double> predictions = new HashMap<>();
        // 对于用户未评分的项目
        for (Item item : items.values()) {
            if (targetUser.getRatings().containsKey(item.getId())) continue;
            // 计算项目与其他已评分项目的相似度
            double totalSimilarity = 0;
            double weightedRating = 0;
            for (Map.Entry<Integer, Double> ratedItem : targetUser.getRatings().entrySet()) {
                int ratedItemId = ratedItem.getKey();
                double rating = ratedItem.getValue();
                double similarity = calculateItemSimilarity(item.getId(), ratedItemId);
                totalSimilarity += similarity;
                weightedRating += similarity * rating;
            }
            if (totalSimilarity > 0) {
                predictions.put(item.getId(), weightedRating / totalSimilarity);
            }
        }
        return getTopN(predictions, topN);
    }
    /**
     * 计算用户相似度(皮尔逊相关系数)
     */
    private double calculateUserSimilarity(User user1, User user2) {
        Set<Integer> commonItems = new HashSet<>(user1.getRatings().keySet());
        commonItems.retainAll(user2.getRatings().keySet());
        if (commonItems.isEmpty()) return 0;
        double mean1 = user1.getRatings().values().stream()
            .mapToDouble(Double::doubleValue).average().orElse(0);
        double mean2 = user2.getRatings().values().stream()
            .mapToDouble(Double::doubleValue).average().orElse(0);
        double numerator = 0;
        double denominator1 = 0;
        double denominator2 = 0;
        for (int itemId : commonItems) {
            double r1 = user1.getRatings().get(itemId) - mean1;
            double r2 = user2.getRatings().get(itemId) - mean2;
            numerator += r1 * r2;
            denominator1 += r1 * r1;
            denominator2 += r2 * r2;
        }
        if (denominator1 == 0 || denominator2 == 0) return 0;
        return numerator / (Math.sqrt(denominator1) * Math.sqrt(denominator2));
    }
    /**
     * 计算项目相似度(余弦相似度)
     */
    private double calculateItemSimilarity(int itemId1, int itemId2) {
        Item item1 = items.get(itemId1);
        Item item2 = items.get(itemId2);
        if (item1 == null || item2 == null) return 0;
        Set<String> commonFeatures = new HashSet<>(item1.getFeatures().keySet());
        commonFeatures.retainAll(item2.getFeatures().keySet());
        if (commonFeatures.isEmpty()) return 0;
        double dotProduct = 0;
        double norm1 = 0;
        double norm2 = 0;
        for (String feature : commonFeatures) {
            double v1 = item1.getFeatures().get(feature);
            double v2 = item2.getFeatures().get(feature);
            dotProduct += v1 * v2;
            norm1 += v1 * v1;
            norm2 += v2 * v2;
        }
        if (norm1 == 0 || norm2 == 0) return 0;
        return dotProduct / (Math.sqrt(norm1) * Math.sqrt(norm2));
    }
    /**
     * 获取Top N推荐
     */
    private Map<Integer, Double> getTopN(Map<Integer, Double> predictions, int n) {
        List<Map.Entry<Integer, Double>> sorted = 
            new ArrayList<>(predictions.entrySet());
        sorted.sort((a, b) -> Double.compare(b.getValue(), a.getValue()));
        Map<Integer, Double> topN = new LinkedHashMap<>();
        int count = Math.min(n, sorted.size());
        for (int i = 0; i < count; i++) {
            topN.put(sorted.get(i).getKey(), sorted.get(i).getValue());
        }
        return topN;
    }
}

的推荐

// ContentBasedFiltering.java
package com.recommend.algorithm;
import com.recommend.model.User;
import com.recommend.model.Item;
import java.util.*;
public class ContentBasedFiltering {
    private Map<Integer, Item> items;
    private Map<Integer, User> users;
    public ContentBasedFiltering(Map<Integer, User> users, Map<Integer, Item> items) {
        this.users = users;
        this.items = items;
    }
    /**
     * 基于内容的推荐
     */
    public Map<Integer, Double> recommend(int userId, int topN) {
        User user = users.get(userId);
        if (user == null) return Collections.emptyMap();
        // 1. 构建用户偏好向量(基于已评分项目)
        Map<String, Double> userPreference = buildUserPreference(user);
        // 2. 计算未评分项目与用户偏好的相似度
        Map<Integer, Double> scores = new HashMap<>();
        for (Item item : items.values()) {
            if (user.getRatings().containsKey(item.getId())) continue;
            double similarity = calculateSimilarity(userPreference, item.getFeatures());
            scores.put(item.getId(), similarity);
        }
        // 3. 返回Top N
        List<Map.Entry<Integer, Double>> sorted = 
            new ArrayList<>(scores.entrySet());
        sorted.sort((a, b) -> Double.compare(b.getValue(), a.getValue()));
        Map<Integer, Double> topN = new LinkedHashMap<>();
        int count = Math.min(topN, sorted.size());
        for (int i = 0; i < count; i++) {
            topN.put(sorted.get(i).getKey(), sorted.get(i).getValue());
        }
        return topN;
    }
    /**
     * 构建用户偏好向量
     */
    private Map<String, Double> buildUserPreference(User user) {
        Map<String, Double> preference = new HashMap<>();
        for (Map.Entry<Integer, Double> entry : user.getRatings().entrySet()) {
            Item item = items.get(entry.getKey());
            if (item == null) continue;
            double rating = entry.getValue();
            for (Map.Entry<String, Double> feature : item.getFeatures().entrySet()) {
                preference.merge(feature.getKey(), 
                    feature.getValue() * rating, 
                    Double::sum);
            }
        }
        // 归一化
        double norm = 0;
        for (double value : preference.values()) {
            norm += value * value;
        }
        norm = Math.sqrt(norm);
        if (norm > 0) {
            preference.replaceAll((k, v) -> v / norm);
        }
        return preference;
    }
    /**
     * 计算余弦相似度
     */
    private double calculateSimilarity(Map<String, Double> userPref, 
                                     Map<String, Double> itemFeatures) {
        Set<String> commonKeys = new HashSet<>(userPref.keySet());
        commonKeys.retainAll(itemFeatures.keySet());
        if (commonKeys.isEmpty()) return 0;
        double dotProduct = 0;
        double userNorm = 0;
        double itemNorm = 0;
        for (String key : commonKeys) {
            dotProduct += userPref.get(key) * itemFeatures.get(key);
        }
        for (double value : userPref.values()) {
            userNorm += value * value;
        }
        for (double value : itemFeatures.values()) {
            itemNorm += value * value;
        }
        userNorm = Math.sqrt(userNorm);
        itemNorm = Math.sqrt(itemNorm);
        if (userNorm == 0 || itemNorm == 0) return 0;
        return dotProduct / (userNorm * itemNorm);
    }
}

混合推荐系统

// HybridRecommender.java
package com.recommend.algorithm;
import java.util.*;
public class HybridRecommender {
    private CollaborativeFiltering collabFiltering;
    private ContentBasedFiltering contentFiltering;
    private double collabWeight = 0.6; // 协同过滤权重
    private double contentWeight = 0.4; // 内容推荐权重
    public HybridRecommender(CollaborativeFiltering collabFiltering, 
                           ContentBasedFiltering contentFiltering) {
        this.collabFiltering = collabFiltering;
        this.contentFiltering = contentFiltering;
    }
    /**
     * 混合推荐
     */
    public Map<Integer, Double> recommend(int userId, int topN) {
        // 获取两种算法的推荐结果
        Map<Integer, Double> collabResults = collabFiltering.recommendByUser(userId, topN);
        Map<Integer, Double> contentResults = contentFiltering.recommend(userId, topN);
        // 融合推荐结果
        Map<Integer, Double> mergedResults = new HashMap<>();
        // 加入协同过滤结果
        for (Map.Entry<Integer, Double> entry : collabResults.entrySet()) {
            mergedResults.put(entry.getKey(), 
                entry.getValue() * collabWeight);
        }
        // 加入内容推荐结果
        for (Map.Entry<Integer, Double> entry : contentResults.entrySet()) {
            double existingScore = mergedResults.getOrDefault(entry.getKey(), 0.0);
            mergedResults.put(entry.getKey(), 
                existingScore + entry.getValue() * contentWeight);
        }
        // 排序并返回Top N
        List<Map.Entry<Integer, Double>> sorted = 
            new ArrayList<>(mergedResults.entrySet());
        sorted.sort((a, b) -> Double.compare(b.getValue(), a.getValue()));
        Map<Integer, Double> finalResults = new LinkedHashMap<>();
        int count = Math.min(topN, sorted.size());
        for (int i = 0; i < count; i++) {
            finalResults.put(sorted.get(i).getKey(), sorted.get(i).getValue());
        }
        return finalResults;
    }
}

数据加载器

// DataLoader.java
package com.recommend.data;
import com.recommend.model.User;
import com.recommend.model.Item;
import java.io.*;
import java.util.*;
public class DataLoader {
    /**
     * 加载用户数据
     */
    public static Map<Integer, User> loadUsers(String filePath) throws IOException {
        Map<Integer, User> users = new HashMap<>();
        try (BufferedReader br = new BufferedReader(new FileReader(filePath))) {
            String line;
            // 跳过表头
            br.readLine();
            while ((line = br.readLine()) != null) {
                String[] parts = line.split(",");
                if (parts.length >= 2) {
                    int id = Integer.parseInt(parts[0].trim());
                    String name = parts[1].trim();
                    users.put(id, new User(id, name));
                }
            }
        }
        return users;
    }
    /**
     * 加载项目数据
     */
    public static Map<Integer, Item> loadItems(String filePath) throws IOException {
        Map<Integer, Item> items = new HashMap<>();
        try (BufferedReader br = new BufferedReader(new FileReader(filePath))) {
            String line;
            // 跳过表头
            br.readLine();
            while ((line = br.readLine()) != null) {
                String[] parts = line.split(",");
                if (parts.length >= 3) {
                    int id = Integer.parseInt(parts[0].trim());
                    String name = parts[1].trim();
                    String category = parts[2].trim();
                    Item item = new Item(id, name, category);
                    // 添加特征(假设从第4列开始是特征)
                    if (parts.length > 3) {
                        for (int i = 3; i < parts.length; i++) {
                            String[] featureParts = parts[i].split(":");
                            if (featureParts.length == 2) {
                                item.addFeature(featureParts[0].trim(), 
                                    Double.parseDouble(featureParts[1].trim()));
                            }
                        }
                    }
                    items.put(id, item);
                }
            }
        }
        return items;
    }
    /**
     * 加载评分数据
     */
    public static void loadRatings(String filePath, Map<Integer, User> users) throws IOException {
        try (BufferedReader br = new BufferedReader(new FileReader(filePath))) {
            String line;
            // 跳过表头
            br.readLine();
            while ((line = br.readLine()) != null) {
                String[] parts = line.split(",");
                if (parts.length >= 3) {
                    int userId = Integer.parseInt(parts[0].trim());
                    int itemId = Integer.parseInt(parts[1].trim());
                    double rating = Double.parseDouble(parts[2].trim());
                    User user = users.get(userId);
                    if (user != null) {
                        user.addRating(itemId, rating);
                    }
                }
            }
        }
    }
}

主程序

// Main.java
package com.recommend;
import com.recommend.algorithm.*;
import com.recommend.data.DataLoader;
import com.recommend.model.User;
import com.recommend.model.Item;
import java.util.*;
public class Main {
    public static void main(String[] args) {
        try {
            // 1. 加载数据
            String basePath = "data/";
            Map<Integer, User> users = DataLoader.loadUsers(basePath + "users.csv");
            Map<Integer, Item> items = DataLoader.loadItems(basePath + "items.csv");
            DataLoader.loadRatings(basePath + "ratings.csv", users);
            System.out.println("加载数据完成:");
            System.out.println("用户数量:" + users.size());
            System.out.println("项目数量:" + items.size());
            // 2. 初始化算法
            CollaborativeFiltering collabFiltering = 
                new CollaborativeFiltering(users, items);
            ContentBasedFiltering contentFiltering = 
                new ContentBasedFiltering(users, items);
            HybridRecommender hybridRecommender = 
                new HybridRecommender(collabFiltering, contentFiltering);
            // 3. 为用户1进行推荐
            int userId = 1;
            int topN = 5;
            System.out.println("\n为用户 " + userId + " 推荐结果:");
            // 基于用户的协同过滤
            Map<Integer, Double> collabResults = 
                collabFiltering.recommendByUser(userId, topN);
            printResults("协同过滤推荐", collabResults, items);
            // 基于内容的推荐
            Map<Integer, Double> contentResults = 
                contentFiltering.recommend(userId, topN);
            printResults("基于内容推荐", contentResults, items);
            // 混合推荐
            Map<Integer, Double> hybridResults = 
                hybridRecommender.recommend(userId, topN);
            printResults("混合推荐", hybridResults, items);
            // 4. 显示用户历史评分
            System.out.println("\n用户 " + userId + " 的历史评分:");
            User user = users.get(userId);
            for (Map.Entry<Integer, Double> rating : user.getRatings().entrySet()) {
                Item item = items.get(rating.getKey());
                System.out.println("  项目: " + item.getName() + 
                    ", 评分: " + rating.getValue());
            }
        } catch (Exception e) {
            e.printStackTrace();
        }
    }
    private static void printResults(String title, 
                                    Map<Integer, Double> results, 
                                    Map<Integer, Item> items) {
        System.out.println("\n" + title + ":");
        for (Map.Entry<Integer, Double> entry : results.entrySet()) {
            Item item = items.get(entry.getKey());
            if (item != null) {
                System.out.printf("  %-20s 预测评分: %.2f%n", 
                    item.getName(), entry.getValue());
            }
        }
    }
}

示例数据文件

// users.csv
id,name
1,Alice
2,Bob
3,Charlie
4,David
5,Eve
// items.csv
id,name,category,genre:action,genre:comedy,genre:drama,genre:sci-fi
1,Matrix,电影,1.0,0.0,0.0,1.0
2,Inception,电影,1.0,0.0,0.5,1.0
3,The Godfather,电影,0.5,0.0,1.0,0.0
4,The Hangover,电影,0.0,1.0,0.5,0.0
5,Interstellar,电影,0.5,0.0,1.0,1.0
6,Avatar,电影,1.0,0.0,0.5,0.5
7,Toy Story,动画,0.0,0.5,0.0,0.0
8,Frozen,动画,0.0,0.5,0.5,0.0
// ratings.csv
userId,itemId,rating
1,1,5
1,2,4
1,4,3
2,1,4
2,3,5
2,6,4
3,2,5
3,5,4
3,7,3
4,4,5
4,7,4
4,8,4
5,1,3
5,3,4
5,6,5

高级特性扩展

// AdvancedRecommendation.java
package com.recommend.algorithm;
import java.util.*;
import java.util.concurrent.*;
public class AdvancedRecommendation {
    // 使用线程池进行并行计算
    private ExecutorService executor = Executors.newFixedThreadPool(4);
    /**
     * 并行协同过滤推荐
     */
    public Future<Map<Integer, Double>> parallelRecommend(int userId, int topN) {
        return executor.submit(() -> {
            // 并行计算逻辑
            return new HashMap<>();
        });
    }
    /**
     * 实时推荐(基于用户最近行为)
     */
    public Map<Integer, Double> realTimeRecommend(int userId, 
                                                 Map<Integer, Double> recentActions) {
        // 根据最近的点击、浏览等行为实时调整推荐
        Map<Integer, Double> adjustedScores = new HashMap<>();
        for (Map.Entry<Integer, Double> action : recentActions.entrySet()) {
            // 计算行为权重
            double weight = action.getValue() * 0.5;
            adjustedScores.merge(action.getKey(), weight, Double::sum);
        }
        return adjustedScores;
    }
    /**
     * 基于聚类的推荐
     */
    public Map<Integer, Double> clusterBasedRecommend(int userId, int topN) {
        // K-means聚类实现
        return new HashMap<>();
    }
    /**
     * 模型评估
     */
    public double evaluateModel(Map<Integer, Map<Integer, Double>> testData) {
        // 计算RMSE或MAE
        return 0.0;
    }
}

使用方法

  1. 创建数据结构:准备用户、项目和评分数据
  2. 选择算法:根据需求选择协同过滤、内容推荐或混合推荐
  3. 调优参数:调整算法参数(如相似度阈值、权重等)
  4. 评估优化:使用交叉验证评估推荐效果

特点

  • ✅ 多种推荐算法实现
  • ✅ 支持混合推荐
  • ✅ 可扩展性强
  • ✅ 包含实时推荐支持
  • ✅ 合理的架构设计

这个推荐系统提供了完整的协同过滤、内容推荐和混合推荐实现,适合作为学习Java推荐系统的入门案例,您可以根据实际需求进行扩展和优化。

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