本文目录导读:

我来为您提供一个完整的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;
}
}
使用方法
- 创建数据结构:准备用户、项目和评分数据
- 选择算法:根据需求选择协同过滤、内容推荐或混合推荐
- 调优参数:调整算法参数(如相似度阈值、权重等)
- 评估优化:使用交叉验证评估推荐效果
特点
- ✅ 多种推荐算法实现
- ✅ 支持混合推荐
- ✅ 可扩展性强
- ✅ 包含实时推荐支持
- ✅ 合理的架构设计
这个推荐系统提供了完整的协同过滤、内容推荐和混合推荐实现,适合作为学习Java推荐系统的入门案例,您可以根据实际需求进行扩展和优化。