diff --git a/python/sglang/srt/mem_cache/storage/hf3fs/hf3fs_usrbio_client.py b/python/sglang/srt/mem_cache/storage/hf3fs/hf3fs_usrbio_client.py index 253219826..0e4e686c3 100644 --- a/python/sglang/srt/mem_cache/storage/hf3fs/hf3fs_usrbio_client.py +++ b/python/sglang/srt/mem_cache/storage/hf3fs/hf3fs_usrbio_client.py @@ -118,47 +118,67 @@ class Hf3fsUsrBioClient(Hf3fsClient): @rsynchronized() def batch_read(self, offsets: List[int], tensors: List[torch.Tensor]) -> List[int]: self.check(offsets, tensors) - + results = [0] * len(offsets) # prepare current = 0 for offset, tensor in zip(offsets, tensors): size = tensor.numel() * tensor.itemsize - self.ior_r.prepare( - self.iov_r[current : current + size], True, self.file, offset - ) - current += size - + try: + self.ior_r.prepare( + self.iov_r[current : current + size], True, self.file, offset + ) + current += size + except Exception as e: + logger.error(f"Error preparing batch read: {e}") + return results # submit ionum = len(offsets) - resv = self.ior_r.submit().wait( - min_results=ionum, timeout=datetime.timedelta(seconds=self.client_timeout) - ) - + try: + resv = self.ior_r.submit().wait( + min_results=ionum, + timeout=datetime.timedelta(seconds=self.client_timeout), + ) + except Exception as e: + logger.error(f"Error submitting batch read: {e}") + return results # results - hf3fs_utils.read_shm(self.shm_r_tensor, tensors) - results = [res.result for res in resv] + try: + hf3fs_utils.read_shm(self.shm_r_tensor, tensors) + results = [res.result for res in resv] + except Exception as e: + logger.error(f"[Hf3fsUsrBioClient] read_shm failed: {e}", exc_info=True) + return results return results @wsynchronized() def batch_write(self, offsets: List[int], tensors: List[torch.Tensor]) -> List[int]: self.check(offsets, tensors) - + results = [0] * len(offsets) # prepare hf3fs_utils.write_shm(tensors, self.shm_w_tensor) current = 0 for offset, tensor in zip(offsets, tensors): size = tensor.numel() * tensor.itemsize - self.ior_w.prepare( - self.iov_w[current : current + size], False, self.file, offset - ) - current += size + try: + self.ior_w.prepare( + self.iov_w[current : current + size], False, self.file, offset + ) + current += size + except Exception as e: + logger.error(f"Error preparing batch write: {e}") + return results # submit ionum = len(offsets) - resv = self.ior_w.submit().wait( - min_results=ionum, timeout=datetime.timedelta(seconds=self.client_timeout) - ) + try: + resv = self.ior_w.submit().wait( + min_results=ionum, + timeout=datetime.timedelta(seconds=self.client_timeout), + ) + except Exception as e: + logger.error(f"Error submitting batch write: {e}") + return results # results results = [res.result for res in resv] diff --git a/python/sglang/srt/mem_cache/storage/hf3fs/storage_hf3fs.py b/python/sglang/srt/mem_cache/storage/hf3fs/storage_hf3fs.py index 55d34dc62..5a5dd0b4d 100644 --- a/python/sglang/srt/mem_cache/storage/hf3fs/storage_hf3fs.py +++ b/python/sglang/srt/mem_cache/storage/hf3fs/storage_hf3fs.py @@ -347,7 +347,11 @@ class HiCacheHF3FS(HiCacheStorage): values: List[torch.Tensor], ) -> List[bool]: page_indices = self.metadata_client.get_page_indices(self.rank, keys) - + if len(page_indices) != len(keys): + logger.error( + f"[Rank {self.rank}] HiCacheHF3FS get: page_indices length {len(page_indices)} mismatch keys length {len(keys)}." + ) + return [False] * len(keys) batch_indices, file_offsets = [], [] for i, page_index in enumerate(page_indices): if page_index is not None: @@ -402,7 +406,16 @@ class HiCacheHF3FS(HiCacheStorage): indices = self.metadata_client.reserve_and_allocate_page_indices( self.rank, key_with_prefix ) - + if len(indices) != len(keys): + logger.error( + f"[Rank {self.rank}] HiCacheHF3FS batch_get: mismatched lengths {len(indices)} != {len(keys)}" + ) + # free allocated pages + if indices: + self.metadata_client.confirm_write( + self.rank, [], [index[1] for index in indices] + ) + return [False] * len(keys) batch_indices, file_offsets, file_values = [], [], [] pages_to_release = []