1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145
| class OcclusionAwareDMS: """ 完整遮挡感知DMS系统 整合: 人脸检测 → 驾驶员识别 → 遮挡检测 → 视线估计 论文Section 4: 完整管线流程 """ def __init__(self, device='cuda'): self.device = device self.face_detector = DualModalFaceDetector() self.driver_identifier = DriverIdentifier().to(device).eval() self.occlusion_detector = OcclusionDetector().to(device).eval() self.gaze_estimator_rgb = GazeRegionEstimator(input_channels=3).to(device).eval() self.gaze_estimator_ir = GazeRegionEstimator(input_channels=1).to(device).eval() def process_frame(self, rgb_frame: np.ndarray, ir_frame: np.ndarray = None) -> dict: """ 处理单帧图像,输出完整DMS结果 Args: rgb_frame: RGB摄像头帧 (H, W, 3) BGR ir_frame: IR摄像头帧 (H, W) 可选 Returns: result: { 'face_detected': bool, 'driver_id': int, 'occlusion': str, # 'clear'|'partial'|'severe' 'gaze_zone': str, 'gaze_confidence': float, 'system_status': str, # 'normal'|'degraded'|'failed' 'modality': str # 'rgb'|'ir' } """ result = { 'face_detected': False, 'driver_id': -1, 'occlusion': 'unknown', 'gaze_zone': 'unknown', 'gaze_confidence': 0.0, 'system_status': 'failed', 'modality': 'rgb' } det_result = self.face_detector.detect(rgb_frame, ir_frame) if not det_result['faces']: return result result['face_detected'] = True result['modality'] = det_result['modality'] face = max(det_result['faces'], key=lambda f: (f['bbox'][2]-f['bbox'][0]) * (f['bbox'][3]-f['bbox'][1])) x1, y1, x2, y2 = face['bbox'] face_img = rgb_frame[y1:y2, x1:x2] if face_img.size == 0: return result driver_tensor = self._preprocess_face(face_img, size=160) with torch.no_grad(): _, driver_logits = self.driver_identifier(driver_tensor) driver_id = driver_logits.argmax(dim=-1).item() result['driver_id'] = driver_id occ_tensor = self._preprocess_face(face_img, size=224) with torch.no_grad(): occ_logits = self.occlusion_detector(occ_tensor) occ_probs = torch.softmax(occ_logits, dim=-1) occ_pred = occ_probs.argmax(dim=-1).item() occ_labels = ['clear', 'partial', 'severe'] result['occlusion'] = occ_labels[occ_pred] if result['occlusion'] == 'severe': result['system_status'] = 'degraded' return result if result['modality'] == 'rgb': gaze_result = self.gaze_estimator_rgb.predict(face_img, is_ir=False) else: if ir_frame is not None: ir_face = ir_frame[y1:y2, x1:x2] if ir_face.size > 0: gaze_result = self.gaze_estimator_ir.predict(ir_face, is_ir=True) else: result['system_status'] = 'degraded' return result else: result['system_status'] = 'degraded' return result result['gaze_zone'] = gaze_result['zone'] result['gaze_confidence'] = gaze_result['confidence'] result['system_status'] = 'normal' return result def _preprocess_face(self, face_img: np.ndarray, size: int = 224) -> torch.Tensor: """预处理面部图像""" from torchvision import transforms as T transform = T.Compose([ T.ToPILImage(), T.Resize((size, size)), T.ToTensor(), T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) return transform(face_img).unsqueeze(0).to(self.device)
if __name__ == "__main__": import numpy as np dms = OcclusionAwareDMS(device='cpu') rgb_frame = np.random.randint(0, 255, (720, 1280, 3), dtype=np.uint8) ir_frame = np.random.randint(0, 255, (720, 1280), dtype=np.uint8) result = dms.process_frame(rgb_frame, ir_frame) print("=" * 50) print("DMS 系统输出:") print(f" 人脸检测: {result['face_detected']}") print(f" 模态: {result['modality']}") print(f" 驾驶员ID: {result['driver_id']}") print(f" 遮挡等级: {result['occlusion']}") print(f" 视线区域: {result['gaze_zone']}") print(f" 视线置信度: {result['gaze_confidence']:.4f}") print(f" 系统状态: {result['system_status']}") print("=" * 50)
|