diff --git a/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount1.java b/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount1.java index 15c4aae..6344a7f 100644 --- a/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount1.java +++ b/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount1.java @@ -1,12 +1,62 @@ package com.github.hcsp.multithread; +import java.io.BufferedReader; import java.io.File; +import java.io.FileNotFoundException; +import java.io.FileReader; +import java.util.ArrayList; +import java.util.HashMap; import java.util.List; import java.util.Map; +import java.util.concurrent.*; public class MultiThreadWordCount1 { // 使用threadNum个线程,并发统计文件中各单词的数量 - public static Map count(int threadNum, List files) { - return null; + // 使用线程池 + public static Map count(int threadNum, List files) throws FileNotFoundException, ExecutionException, InterruptedException { + List>> resultsFuture = new ArrayList<>(); + ExecutorService threadPool = Executors.newFixedThreadPool(threadNum); + for (File file : files) { + BufferedReader bufferedReader = new BufferedReader(new FileReader(file)); + Future> mapFuture = threadPool.submit(new readFileCallable(bufferedReader)); + resultsFuture.add(mapFuture); + } + threadPool.shutdown(); + return mergeCountResults(resultsFuture); + } + + private static Map mergeCountResults(List>> resultsFuture) throws ExecutionException, InterruptedException { + Map finalResults = new HashMap<>(); + for (Future> mapFuture : resultsFuture) { + Map stringIntegerMap = mapFuture.get(); + for (Map.Entry stringIntegerEntry : stringIntegerMap.entrySet()) { + String word = stringIntegerEntry.getKey(); + int updatedValue = finalResults.getOrDefault(word, 0) + stringIntegerEntry.getValue(); + finalResults.put(word, updatedValue); + } + } + return finalResults; + } + + public static class readFileCallable implements Callable> { + BufferedReader bufferedReader; + + public readFileCallable(BufferedReader bufferedReader) { + this.bufferedReader = bufferedReader; + } + + @Override + public Map call() throws Exception { + Map fileWords = new HashMap<>(); + String line; + while ((line = bufferedReader.readLine()) != null) { + String[] words = line.split(" "); + for (String word : words) { + fileWords.put(word, fileWords.getOrDefault(word, 0) + 1); + } + } + return fileWords; + } } } + diff --git a/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount2.java b/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount2.java index 3f23afa..d1721c6 100644 --- a/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount2.java +++ b/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount2.java @@ -1,8 +1,73 @@ package com.github.hcsp.multithread; +import java.io.BufferedReader; +import java.io.File; +import java.io.FileReader; +import java.io.IOException; +import java.util.List; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ForkJoinPool; +import java.util.concurrent.ForkJoinTask; +import java.util.concurrent.RecursiveTask; + public class MultiThreadWordCount2 { // 使用threadNum个线程,并发统计文件中各单词的数量 - // public static Map count(int threadNum, List files) { - // return null; - // } + // 使用ForkJoinPool + public static Map count(int threadNum, List files) { + ForkJoinPool forkJoinPool = new ForkJoinPool(threadNum); + ForkJoinTask> task = new readFileWordTask(files); + return forkJoinPool.invoke(task); + } + + public static class readFileWordTask extends RecursiveTask> { + private static final int THRESHOLD = 3; + List files; + + public readFileWordTask(List files) { + this.files = files; + } + + @Override + protected Map compute() { + if (files.size() <= THRESHOLD) { + Map results = new ConcurrentHashMap<>(); + for (File file : files) { + try { + BufferedReader bufferedReader = new BufferedReader(new FileReader(file)); + String line; + while ((line = bufferedReader.readLine()) != null) { + String[] words = line.split(" "); + for (String word : words) { + results.put(word, results.getOrDefault(word, 0) + 1); + } + } + // System.out.println(file.getCanonicalFile()); + } catch (IOException e) { + e.printStackTrace(); + } + } + return results; + } else { + List subFiles1 = files.subList(0, files.size() / 2); + List subFiles2 = files.subList(files.size() / 2, files.size()); + readFileWordTask subTask1 = new readFileWordTask(subFiles1); + readFileWordTask subTask2 = new readFileWordTask(subFiles2); + invokeAll(subTask1, subTask2); + Map subResult1 = subTask1.join(); + Map subResult2 = subTask2.join(); + return mergeTwoMap(subResult1, subResult2); + } + } + + private Map mergeTwoMap(Map subResult1, Map subResult2) { + Map results = new ConcurrentHashMap<>(subResult2); + for (Map.Entry stringIntegerEntry : subResult1.entrySet()) { + String word = stringIntegerEntry.getKey(); + int countNum = results.getOrDefault(word, 0) + stringIntegerEntry.getValue(); + results.put(word, countNum); + } + return results; + } + } } diff --git a/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount3.java b/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount3.java index e180ee2..152570a 100644 --- a/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount3.java +++ b/src/main/java/com/github/hcsp/multithread/MultiThreadWordCount3.java @@ -1,8 +1,77 @@ package com.github.hcsp.multithread; +import java.io.*; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; + + public class MultiThreadWordCount3 { // 使用threadNum个线程,并发统计文件中各单词的数量 - // public static Map count(int threadNum, List files) { - // return null; - // } + // 使用CountDownLatch + + public static Map count(int threadNum, List files) throws FileNotFoundException { + CountDownLatch countDownLatch = new CountDownLatch(files.size()); + List> mapList = new CopyOnWriteArrayList<>(); + Map results = new HashMap<>(); + ExecutorService pool = Executors.newFixedThreadPool(threadNum); + for (File file : files) { + BufferedReader bufferedReader = new BufferedReader(new FileReader(file)); + pool.execute(new ReadFileRunnable(bufferedReader, mapList, countDownLatch)); + } + try { + countDownLatch.await(); + for (Map map : mapList) { + for (Map.Entry stringIntegerEntry : map.entrySet()) { + String word = stringIntegerEntry.getKey(); + int countNum = results.getOrDefault(word, 0) + stringIntegerEntry.getValue(); + results.put(word, countNum); + } + } + } catch (InterruptedException e) { + e.printStackTrace(); + } finally { + pool.shutdown(); + } + return results; + } + + public static class ReadFileRunnable implements Runnable { + BufferedReader bufferedReader; + List> mapList; + CountDownLatch latch; + + public ReadFileRunnable(BufferedReader bufferedReader, List> mapList, CountDownLatch latch) { + this.bufferedReader = bufferedReader; + this.mapList = mapList; + this.latch = latch; + } + + @Override + public void run() { + Map fileResults = new HashMap<>(); + String line = null; + while (true) { + try { + if ((line = bufferedReader.readLine()) == null) { + break; + } + } catch (IOException e) { + e.printStackTrace(); + } + assert line != null; + String[] words = line.split(" "); + for (String word : words) { + fileResults.put(word, fileResults.getOrDefault(word, 0) + 1); + } + } + this.latch.countDown(); + mapList.add(fileResults); + } + } } +