
這篇文章會用 TDD 手刻 mySumOf、myMaxByOrNull、myMinByOrNull、myAverage,理解這些函式和 fold/reduce 的關係,開始進入聚合篇
| Kotlin | C# LINQ | 備註 |
|---|---|---|
sumOf { it.salary } |
Sum(x => x.Salary) |
|
maxByOrNull { it.salary } |
MaxBy(x => x.Salary) (.NET 6+) |
|
minByOrNull { it.salary } |
MinBy(x => x.Salary) (.NET 6+) |
|
maxOfOrNull { it.salary } |
Max(x => x.Salary) |
maxOf 回傳值,maxBy 回傳元素 |
average() |
Average() |
這三個名字很像但行為不同,容易搞混
maxByOrNull { it.salary } — 回傳整個元素。「薪水最高的那個 Employee 是誰?」回傳 Grace(Employee 物件)
maxOfOrNull { it.salary } — 回傳 selector 的結果。「最高薪水是多少?」回傳 92000(Int)
max() — 直接比較元素本身,需要 T : Comparable<T>。「這些數字裡最大的?」對 List<Int> 有用,對 List<Employee> 就無法呼叫
employees.maxByOrNull { it.salary } // Grace (Employee)
employees.maxOfOrNull { it.salary } // 92000 (Int)
listOf(3, 7, 2).max() // 7 (Int)
@Test
fun `sumOf employee salaries`() {
val result = employees.mySumOf { it.salary }
assertEquals(507000, result)
}
@Test
fun `sumOf empty list returns zero`() {
val empty = emptyList<Employee>()
val result = empty.mySumOf { it.salary }
assertEquals(0, result)
}
@Test
fun `sumOf with transform`() {
val numbers = listOf(1, 2, 3, 4, 5)
val result = numbers.mySumOf { it * it }
assertEquals(55, result)
}
空集合回傳 0,不是拋例外。加法的初始值天生就是 0
inline fun <T> Iterable<T>.mySumOf(selector: (T) -> Int): Int {
var sum = 0
for (element in this) {
sum += selector(element)
}
return sum
}
就是 fold(0) { acc, element -> acc + selector(element) } 的展開版本。用 for 迴圈寫語意更直接,少一層 fold 的包裝(fold 本身是 inline 函式,Lambda 會被內聯,所以兩種寫法的執行成本其實一樣,差別只在讀起來的清楚程度)
注意回傳型別寫死了 Int。stdlib 的 sumOf 有好幾個多載版本:回傳 Int、Long、Double、UInt 等等。每種數值型別一個,因為 Kotlin 的泛型沒辦法表達「任何支援加法的型別」。我們在這裡只先實作 Int 版本
這個 for 迴圈累加的版本已經跟 stdlib 的 sumOf 同一個形狀,沒什麼好改的。真正要追問的是「為什麼回傳型別要寫死」,下一節再來談
看到 stdlib 一堆 sumOf 多載你可能會想:寫一個 sumOf<T : Number> 不就好了
問題在 Number 沒有 + 運算子
Java 的 Number 是抽象類別,只有 intValue()、doubleValue() 之類的取值方法,沒有運算行為。Kotlin 跑在 JVM 上繼承這個限制,沒有「支援加法的型別」這種泛型約束可寫
C# 走過同樣的歷程。LINQ 的 Sum 從 .NET 3.5 引入到 .NET 6,公開 API 一直都是針對每個數值型別各寫一個 overload。直到 C# 11 / .NET 7 加上 generic math(INumber<T>),才有辦法自己寫出一個泛型版的數值聚合
要注意的是,LINQ 內建的 Enumerable.Sum 至今仍維持各型別 overload 的公開簽名,並沒有改成泛型 API。下面這段是「有了 generic math 之後,你可以自己這樣寫」的示意,不是 LINQ 的內建簽名
// 自己用 generic math 寫的版本,非 LINQ 內建 API
public static T MySum<T>(this IEnumerable<T> source) where T : INumber<T> {
T sum = T.Zero;
foreach (var item in source) sum += item;
return sum;
}
INumber<T> 定義了 +、-、Zero、One 等 operator,讓泛型型別可以做數值運算。這需要 C# 11 的 static abstract interface members 的配合
Kotlin 暫時沒有這層抽象。stdlib 的 sumOf 是針對每個型別寫獨立 overload,呼叫端只看到「sumOf」一個名字,內部分支由 compiler 解析。代價是 stdlib 內部要為每種數值型別個別實作
所以 mySumOf 寫死 Int 不是偷懶,是 Kotlin 在 JVM 上的型別系統限制。要嘛跟 stdlib 一樣做手工 overload,要嘛等 Kotlin 加上類似 INumber<T> 的 generic math 介面(目前還沒)
@Test
fun `maxByOrNull employee salary`() {
val result = employees.myMaxByOrNull { it.salary }
assertEquals("Grace", result?.name)
}
@Test
fun `maxByOrNull empty list returns null`() {
val empty = emptyList<Employee>()
val result = empty.myMaxByOrNull { it.salary }
assertNull(result)
}
@Test
fun `maxByOrNull string by length`() {
val words = listOf("fig", "banana", "kiwi", "apple")
val result = words.myMaxByOrNull { it.length }
assertEquals("banana", result)
}
@Test
fun `minByOrNull employee salary`() {
val result = employees.myMinByOrNull { it.salary }
assertEquals("Frank", result?.name)
}
@Test
fun `minByOrNull empty list returns null`() {
val empty = emptyList<Employee>()
val result = empty.myMinByOrNull { it.salary }
assertNull(result)
}
回傳的是元素本身(Employee 或 String),不是 selector 的結果
inline fun <T, R : Comparable<R>> Iterable<T>.myMaxByOrNull(selector: (T) -> R): T? {
val iterator = this.iterator()
if (!iterator.hasNext()) {
return null
}
var maxElement = iterator.next()
var maxValue = selector(maxElement)
while (iterator.hasNext()) {
val element = iterator.next()
val value = selector(element)
if (value > maxValue) {
maxElement = element
maxValue = value
}
}
return maxElement
}
需要追蹤兩樣東西:目前最大的元素(maxElement)和對應的 selector 結果(maxValue)。每次遇到更大的 value 就更新兩者
泛型簽名是 <T, R : Comparable<R>>。T 是元素型別(不需要 Comparable),R 是 selector 回傳的型別(需要 Comparable,因為要比大小)。跟 day 15 sortedBy 的約束一樣
myMinByOrNull 跟它只差一個比較符號:> 變 <
inline fun <T, R : Comparable<R>> Iterable<T>.myMinByOrNull(selector: (T) -> R): T? {
val iterator = this.iterator()
if (!iterator.hasNext()) {
return null
}
var minElement = iterator.next()
var minValue = selector(minElement)
while (iterator.hasNext()) {
val element = iterator.next()
val value = selector(element)
if (value < minValue) {
minElement = element
minValue = value
}
}
return minElement
}
兩個函式的實作已經是 stdlib 的形狀,沒什麼好改的。stdlib 多做的細節留到後面「與 stdlib 原始碼比較」一節,跟 maxByOrNull 的原始碼一起看
@Test
fun `average of integers`() {
val numbers = listOf(10, 20, 30, 40, 50)
val result = numbers.myAverage()
assertEquals(30.0, result)
}
@Test
fun `average of empty list returns NaN`() {
val empty = emptyList<Int>()
val result = empty.myAverage()
assertTrue(result.isNaN())
}
空集合回傳 Double.NaN(Not a Number),不是 0.0 也不是拋例外。stdlib 就是這樣設計的:0 / 0 在數學上沒有意義,用 NaN 表示
fun Iterable<Int>.myAverage(): Double {
var sum = 0.0
var count = 0
for (element in this) {
sum += element
count++
}
return if (count == 0) Double.NaN else sum / count
}
注意 receiver 型別是 Iterable<Int> 而不是泛型的 Iterable<T>。跟 sumOf 一樣的問題:沒辦法用泛型表達「任何數值型別」,所以 stdlib 針對 Int、Long、Double 等各寫一個版本
sum 用 0.0(Double) 而不是 0(Int),避免整數除法。10 / 3 在 Kotlin 裡是 3(整數除法),10.0 / 3 才是 3.333...
累加、計數、最後相除,這已經是 stdlib average 的骨架,空集合回傳 NaN 的處理也一樣,本篇就停在這個版本。stdlib 在細節上多做的事,下一節拿 maxByOrNull 的原始碼來看
原始碼位置:kotlin.collections 的 _Collections.kt
stdlib 的 maxByOrNull 寫法跟我們的幾乎一樣
public inline fun <T, R : Comparable<R>> Iterable<T>.maxByOrNull(
selector: (T) -> R
): T? {
val iterator = iterator()
if (!iterator.hasNext()) return null
var maxElem = iterator.next()
if (!iterator.hasNext()) return maxElem
var maxValue = selector(maxElem)
do {
val e = iterator.next()
val v = selector(e)
if (maxValue < v) {
maxElem = e
maxValue = v
}
} while (iterator.hasNext())
return maxElem
}
stdlib 多了一個最佳化:如果只有一個元素,直接回傳不呼叫 selector。然後用 do-while 而不是 while,因為進入迴圈前已經確認 hasNext() 為 true
另一個細節:stdlib 用 maxValue < v 而不是 v > maxValue。效果一樣,但 maxValue < v 把被比較的值放前面,風格上更一致
所有數值聚合都能用 fold 改寫
// sumOf 用 fold
employees.fold(0) { acc, emp -> acc + emp.salary }
// maxByOrNull 用 fold(但要處理空集合)
employees.fold(null as Employee?) { acc, emp ->
if (acc == null || emp.salary > acc.salary) emp else acc
}
能用 fold 寫不代表一定要用 fold 寫。sumOf { it.salary } 一看就懂,fold(0) { acc, emp -> acc + emp.salary } 要多想一下。專用函式存在的理由是可讀性
上面講的是呼叫時該挑哪個 API,不過在實作時又是另一件事。day 19 的 myFold 和 myReduce 都是一行委派給 myAggregate,這篇的四個函式卻全是手寫迴圈,看起來不一致
mySumOf 是可以委派的那一個。改寫成 fold(0) { acc, e -> acc + selector(e) } 行為一樣、成本也一樣(前面提過 fold 是 inline),選 for 迴圈純粹是可讀性
myMaxByOrNull 委派之後,selector 會被多算一輪。上面那個 fold 版裡,emp.salary 和 acc.salary 各算一次,泛型化之後這兩個讀取就是 selector(emp) 和 selector(acc),所以除了第一個元素(|| 短路掉了),每個元素都要算兩次;手寫版把 selector 的結果存在 maxValue 這個變數裡,只在真的換人時才更新,每個元素只算一次。selector 是 it.salary 這種欄位讀取時差別看不出來,換成字串解析或查表就是兩倍成本
還有一點,fold 版的累加器被迫宣告成 T?,等於拿 null 當「還沒開始」的哨兵。元素型別本身允許 null 的時候(List<Employee?>),這個哨兵就分不出「還沒開始」和「這個元素就是 null」,跟 day 19 aggregate 是同一類問題,差別在那邊有 map 可以問 containsKey,這裡只能再開一個 boolean。手寫版用 hasNext() 判斷開頭,從頭就沒這個歧義
myMaxByOrNull 改用 reduce 也不行,空集合會拋例外(day 18 談過),換成 reduceOrNull 合約是補回來了,但每個元素照樣算兩次 selector,省不掉
myAverage 卡在別的地方:它要同時累積 sum 和 count 兩樣東西,fold 只有一個累加器,得把兩個值包成 Pair 或 data class 一路傳下去,每跑一個元素就多配置一個物件。用兩個區域變數就沒這回事。回頭看 myMaxByOrNull,想省掉多算的那一次,也只能把 (element, value) 包成 Pair 當累加器,計算是省下來了,換來的就是這裡的配置
最直接的證據是 stdlib 自己也不是走委派。前面貼的 maxByOrNull 原始碼是手寫的 do-while,_Collections.kt 裡的 sumOf 和 average 也是同一套寫法,沒有一個轉呼叫 fold。「這個函式可以用 fold 表達」跟「這個函式該用 fold 實作」要分清楚
這篇的函式都是 fold/reduce 的語法糖,寫起來比 fold 簡單,讀起來也更清楚。專用函式不是邏輯多厲害,而是讓人一眼看出意圖
下一篇講 joinToString,把集合轉成字串。也是一種聚合,但目標型別固定是 String,而且多了分隔符號、前綴、後綴這些參數
同步刊登於 Blog
圖片來源:AI 產生