有限个已排序无限流合并为单一已排序无限流的实现方法
Great question! Merging sorted infinite streams is a classic problem that plays perfectly to the strengths of lazy streams—since we never need to generate all elements upfront, just compute them on demand. Let's break down how to do this properly, with Scala code examples matching your function signature.
Core Idea
The key here is lazy evaluation and always selecting the smallest available next element from all the input streams. Since the streams are infinite and sorted, we never have to worry about exhausting all elements (though our code can still handle finite streams for robustness). We'll build the merged stream one element at a time, updating our set of active streams each time we take an element.
Naive Implementation (Simple, O(n) per element)
This approach is straightforward: for each step, we find the smallest head element across all non-empty streams, then recursively merge the remaining streams (with the chosen stream's tail replacing its head).
def merge[T](ss: List[Stream[T]])(implicit ord: Ordering[T]): Stream[T] = { // Filter out any exhausted streams (handles finite streams too) val nonEmptyStreams = ss.filter(_.nonEmpty) nonEmptyStreams match { case Nil => Stream.empty[T] case _ => // Find the smallest head element among all active streams val minHead = nonEmptyStreams.map(_.head).min(ord) // For each stream: if its head is the min, replace it with its tail; keep others as-is val nextStreams = nonEmptyStreams.flatMap { stream => if (ord.equiv(stream.head, minHead)) stream.tail :: Nil else stream :: Nil } // Lazily prepend the min element to the result of merging the next set of streams minHead #:: merge(nextStreams) } }
Notes on the naive approach:
- We use
Ordering[T]as an implicit parameter to support any sortable type (Int, String, custom types with an Ordering instance). - The
#::operator is Scala's lazy stream constructor—it doesn't evaluate the recursivemerge(nextStreams)call until someone asks for the next element in the merged stream. - This works perfectly for infinite streams, but each element retrieval takes O(n) time (where n is the number of input streams), which can be slow if n is large.
Optimized Implementation (Priority Queue, O(log n) per element)
To speed things up, we can use a min-heap (priority queue) to keep track of the next available element from each stream. This reduces the time per element retrieval to O(log n), which is much better for larger numbers of input streams.
import scala.collection.mutable.PriorityQueue def merge[T](ss: List[Stream[T]])(implicit ord: Ordering[T]): Stream[T] = { // Scala's PriorityQueue is a max-heap by default, so we reverse the ordering to make it a min-heap val reverseOrdering = ord.reverse // Initialize the heap with non-empty streams, storing (current head, remaining stream) val heap = PriorityQueue.empty[(T, Stream[T])](reverseOrdering) ss.filter(_.nonEmpty).foreach { stream => heap.enqueue((stream.head, stream.tail)) } // Recursive helper to build the merged stream def loop(): Stream[T] = { if (heap.isEmpty) Stream.empty[T] else { val (minValue, restOfStream) = heap.dequeue() // If the stream still has elements, add its next head and tail back to the heap if (restOfStream.nonEmpty) { heap.enqueue((restOfStream.head, restOfStream.tail)) } // Lazily prepend the min value to the next iteration of the loop minValue #:: loop() } } loop() }
Notes on the optimized approach:
- The heap maintains the smallest available element at the top, so we always get the next element in O(1) time (plus O(log n) time to rebalance the heap after dequeuing/enqueuing).
- We only ever store the next element from each active stream in the heap, keeping memory usage efficient even with many input streams.
- Like the naive approach, this is fully lazy—elements are only computed when requested.
Example Usage
Let's test this with a few infinite sorted streams:
// Create infinite sorted streams val evenNumbers = Stream.from(0, 2) // 0, 2, 4, 6, ... val oddNumbers = Stream.from(1, 2) // 1, 3, 5, 7, ... val multiplesOf3 = Stream.from(0, 3) // 0, 3, 6, 9, ... // Merge them val mergedStream = merge(List(evenNumbers, oddNumbers, multiplesOf3)) // Take the first 10 elements to verify mergedStream.take(10).foreach(println) // Output: 0, 0, 1, 2, 3, 3, 4, 6, 6, 7
内容的提问来源于stack exchange,提问作者ntviet18

