1. spark cache原理
Task运行的时候是要去获取Parent的RDD对应的Partition的数据的,即它会调用RDD的iterator方法把对应的Partition的数据集给遍历出来,具体流程如下图:

从图中可以看出,spark cache的本质就是将RDD的数据存储在了BlockManager上,下次重新使用的时候直接从BlockManager获取即可,免去了从“头”计算的开销。
2.cache 源代码分析
首先还是从RDD.scala的iterator方法开始,如果storageLevel不等于None,则调用getOrCompute,如果storageLevel等于None,则调用computeOrReadCheckpoint从头开始计算或者从checkpoint读取。
1final def iterator(split: Partition, context: TaskContext): Iterator[T] = { 2 // storageLevel不等于NONE,说明RDD已经cache 3 if (storageLevel != StorageLevel.NONE) { 4 getOrCompute(split, context) 5 } else { 6 // 进行rdd partition的计算或者从checkpoint读取数据 7 computeOrReadCheckpoint(split, context) 8 } 9}
getOrCompute方法中会调用BlockManager的getOrElseUpdate方法,如果指定的block存在,则直接获取,否则调用computeOrReadCheckpoint方法去计算block,然后再保存到BlockManager。
1private[spark] def getOrCompute(partition: Partition, context: TaskContext): Iterator[T] = { 2 val blockId = RDDBlockId(id, partition.index) 3 var readCachedBlock = true 4 SparkEnv.get.blockManager.getOrElseUpdate(blockId, storageLevel, elementClassTag, () => { 5 readCachedBlock = false 6 computeOrReadCheckpoint(partition, context) 7 }) match { 8 case Left(blockResult) => 9 if (readCachedBlock) { 10 // 如果已经被缓存则直接读取 11 val existingMetrics = context.taskMetrics().inputMetrics 12 existingMetrics.incBytesRead(blockResult.bytes) 13 new InterruptibleIterator[T](context, blockResult.data.asInstanceOf[Iterator[T]]) { 14 override def next(): T = { 15 existingMetrics.incRecordsRead(1) 16 delegate.next() 17 } 18 } 19 } else { 20 new InterruptibleIterator(context, blockResult.data.asInstanceOf[Iterator[T]]) 21 } 22 case Right(iter) => 23 new InterruptibleIterator(context, iter.asInstanceOf[Iterator[T]]) 24 } 25 } 26 27 28def getOrElseUpdate[T]( 29 blockId: BlockId, 30 level: StorageLevel, 31 classTag: ClassTag[T], 32 makeIterator: () => Iterator[T]): Either[BlockResult, Iterator[T]] = { 33 // 尝试从本地获取数据,如果获取不到则从远端获取 34 get[T](blockId)(classTag) match { 35 case Some(block) => 36 return Left(block) 37 case _ => 38 } 39 // 如果本地化和远端都没有获取到数据,则调用makeIterator计算,最后将结果写入block 40 doPutIterator(blockId, makeIterator, level, classTag, keepReadLock = true) match { 41 case None => 42 val blockResult = getLocalValues(blockId).getOrElse { 43 releaseLock(blockId) 44 throw new SparkException(s"get() failed for block $blockId even though we held a lock") 45 } 46 releaseLock(blockId) 47 Left(blockResult) 48 case Some(iter) => 49 Right(iter) 50 } 51 }
computeOrReadCheckpoint方法中会判断rdd是否checkpoint,如果有则调用第一个parent rdd的iterator方法获取,否则从“头”开始计算。
1private[spark] def computeOrReadCheckpoint(split: Partition, context: TaskContext): Iterator[T] = 2 { 3 if (isCheckpointedAndMaterialized) { 4 //如果rdd被checkpointed,则调用第一个parent rdd的iterator方法获取 5 firstParent[T].iterator(split, context) 6 } else { 7 //如果rdd没被checkpointed,则重新计算 8 compute(split, context) 9 } 10 }