diff --git a/src/enhance_pedestrian_safety/annotation_generator.py b/src/enhance_pedestrian_safety/annotation_generator.py index d443854d96..734fa91616 100644 --- a/src/enhance_pedestrian_safety/annotation_generator.py +++ b/src/enhance_pedestrian_safety/annotation_generator.py @@ -21,7 +21,12 @@ def detect_objects(self, world, frame_num, timestamp): 'frame_id': frame_num, 'timestamp': timestamp, 'objects': [], - 'camera_info': {} + 'camera_info': {}, + 'safety_info': { + 'pedestrian_count': 0, + 'vehicle_count': 0, + 'high_risk_interactions': 0 + } } try: @@ -36,6 +41,16 @@ def detect_objects(self, world, frame_num, timestamp): annotations['objects'].append(obj_info) self.object_counter += 1 + # 更新安全统计 + if 'walker' in obj_type: + annotations['safety_info']['pedestrian_count'] += 1 + elif 'vehicle' in obj_type: + annotations['safety_info']['vehicle_count'] += 1 + + # 检测高风险交互 + annotations['safety_info']['high_risk_interactions'] = self._detect_high_risk_interactions( + annotations['objects']) + self._save_annotations(frame_num, annotations) return annotations @@ -110,6 +125,30 @@ def _get_object_class(self, type_id): else: return 'unknown' + def _detect_high_risk_interactions(self, objects): + """检测高风险交互""" + high_risk_count = 0 + pedestrians = [obj for obj in objects if obj['class'] == 'pedestrian'] + vehicles = [obj for obj in objects if obj['class'] in ['car', 'vehicle']] + + for pedestrian in pedestrians: + for vehicle in vehicles: + # 计算距离 + p_loc = pedestrian['location'] + v_loc = vehicle['location'] + + distance = np.sqrt( + (p_loc['x'] - v_loc['x']) ** 2 + + (p_loc['y'] - v_loc['y']) ** 2 + + (p_loc['z'] - v_loc['z']) ** 2 + ) + + # 如果距离小于5米,认为是高风险交互 + if distance < 5.0: + high_risk_count += 1 + + return high_risk_count + def _save_annotations(self, frame_num, annotations): """保存标注到文件""" filename = f"frame_{frame_num:06d}.json" @@ -130,7 +169,13 @@ def _update_master_annotation(self): 'total_frames': len(self.frame_annotations), 'total_objects': self.object_counter, 'frames': list(self.frame_annotations.keys()), - 'created': datetime.now().isoformat() + 'created': datetime.now().isoformat(), + 'safety_summary': { + 'total_pedestrians': sum(f['safety_info']['pedestrian_count'] for f in self.frame_annotations.values()), + 'total_vehicles': sum(f['safety_info']['vehicle_count'] for f in self.frame_annotations.values()), + 'total_high_risk': sum( + f['safety_info']['high_risk_interactions'] for f in self.frame_annotations.values()) + } } with open(master_file, 'w', encoding='utf-8') as f: @@ -141,6 +186,7 @@ def generate_summary(self): vehicle_count = 0 pedestrian_count = 0 other_count = 0 + high_risk_count = 0 for frame_data in self.frame_annotations.values(): for obj in frame_data.get('objects', []): @@ -151,14 +197,21 @@ def generate_summary(self): else: other_count += 1 + high_risk_count += frame_data['safety_info']['high_risk_interactions'] + summary = { 'total_frames': len(self.frame_annotations), 'total_objects': self.object_counter, 'vehicles': vehicle_count, 'pedestrians': pedestrian_count, 'other_objects': other_count, + 'high_risk_interactions': high_risk_count, 'average_objects_per_frame': self.object_counter / len( - self.frame_annotations) if self.frame_annotations else 0 + self.frame_annotations) if self.frame_annotations else 0, + 'safety_metrics': { + 'pedestrian_to_vehicle_ratio': pedestrian_count / max(1, vehicle_count), + 'high_risk_percentage': high_risk_count / max(1, self.object_counter) * 100 + } } summary_file = os.path.join(self.annotations_dir, "summary.json") diff --git a/src/enhance_pedestrian_safety/carla_utils.py b/src/enhance_pedestrian_safety/carla_utils.py index 0fd30ef025..349f4e7c78 100644 --- a/src/enhance_pedestrian_safety/carla_utils.py +++ b/src/enhance_pedestrian_safety/carla_utils.py @@ -1,7 +1,6 @@ """ CARLA 工具模块 - 用于自动查找和配置 CARLA 路径 """ -# carla_utils.py import sys import os import glob diff --git a/src/enhance_pedestrian_safety/config_manager.py b/src/enhance_pedestrian_safety/config_manager.py index aa07724783..7276b09a15 100644 --- a/src/enhance_pedestrian_safety/config_manager.py +++ b/src/enhance_pedestrian_safety/config_manager.py @@ -6,6 +6,7 @@ try: import yaml + YAML_AVAILABLE = True except ImportError: YAML_AVAILABLE = False @@ -65,6 +66,13 @@ def suggest_optimizations(config: Dict[str, Any]) -> List[str]: if len(enabled_outputs) > 5: suggestions.append(f"启用的输出类型过多({len(enabled_outputs)}),可能影响性能,建议只启用必要的输出") + # 行人安全相关建议 + if config.get('traffic', {}).get('pedestrians', 0) < 5: + suggestions.append("行人数量较少,建议增加行人数量以更好地测试行人安全") + + if not config.get('v2x', {}).get('enabled', False): + suggestions.append("V2X通信未启用,建议启用以支持行人安全预警") + return suggestions @@ -193,7 +201,8 @@ def optimize_for_safety(config: Dict[str, Any]) -> Dict[str, Any]: 'walker.pedestrian.0002', 'walker.pedestrian.0003', 'walker.pedestrian.0004' - ] + ], + 'speed_limit': 30.0 # 添加车速限制 }) # 优化传感器配置以更好地检测行人 @@ -225,7 +234,9 @@ def optimize_for_safety(config: Dict[str, Any]) -> Dict[str, Any]: v2x.update({ 'enabled': True, 'communication_range': 300.0, - 'update_interval': 1.0 # 更频繁地更新 + 'update_interval': 1.0, # 更频繁地更新 + 'enable_safety_warnings': True, + 'pedestrian_warning_threshold': 10.0 # 行人警告距离阈值 }) coop = optimized.setdefault('cooperative', {}) @@ -233,6 +244,7 @@ def optimize_for_safety(config: Dict[str, Any]) -> Dict[str, Any]: 'num_coop_vehicles': 2, 'enable_shared_perception': True, 'enable_traffic_warnings': True, + 'enable_pedestrian_warnings': True, # 启用行人警告 'enable_maneuver_coordination': False, 'data_fusion_interval': 0.5, # 更频繁地融合 'max_shared_objects': 100, @@ -247,7 +259,8 @@ def optimize_for_safety(config: Dict[str, Any]) -> Dict[str, Any]: 'compression_level': 3, 'enable_memory_cache': True, 'max_cache_size': 40, - 'frame_rate_limit': 8.0 + 'frame_rate_limit': 8.0, + 'safety_monitoring_interval': 1.0 # 安全监控间隔 }) # 输出配置 @@ -260,16 +273,33 @@ def optimize_for_safety(config: Dict[str, Any]) -> Dict[str, Any]: 'save_fusion': True, 'save_cooperative': True, 'save_enhanced': True, + 'save_safety_reports': True, # 保存安全报告 'validate_data': True, 'run_analysis': True, - 'run_quality_check': True + 'run_quality_check': True, + 'generate_safety_summary': True # 生成安全摘要 + }) + + # 增强配置 + enhanced = optimized.setdefault('enhancement', {}) + enhanced.update({ + 'enabled': True, + 'enable_random': True, + 'quality_check': True, + 'save_original': True, + 'save_enhanced': True, + 'calibration_generation': True, + 'enhanced_dir_name': 'enhanced', + 'methods': ['normalize', 'contrast', 'brightness', 'pedestrian_highlight', 'safety_warning'], + 'weather_effects': True, + 'augmentation_level': 'medium', + 'pedestrian_safety_mode': True # 启用行人安全模式 }) return optimized class ConfigManager: - PRESET_CONFIGS = { 'balanced': { 'description': '平衡配置 - 兼顾性能和质量', @@ -338,8 +368,8 @@ def load_config(config_file: Optional[str] = None, preset: Optional[str] = None) def _get_default_config() -> Dict[str, Any]: return { 'scenario': { - 'name': 'multi_sensor_scene', - 'description': '多传感器协同数据采集场景', + 'name': 'pedestrian_safety', + 'description': '行人安全增强数据采集场景', 'town': 'Town10HD', 'weather': 'clear', 'time_of_day': 'noon', @@ -351,7 +381,7 @@ def _get_default_config() -> Dict[str, Any]: 'traffic': { 'ego_vehicles': 1, 'background_vehicles': 8, - 'pedestrians': 6, + 'pedestrians': 12, # 增加默认行人数量 'traffic_lights': True, 'batch_spawn': True, 'max_spawn_attempts': 5, @@ -363,8 +393,11 @@ def _get_default_config() -> Dict[str, Any]: ], 'pedestrian_types': [ 'walker.pedestrian.0001', - 'walker.pedestrian.0002' - ] + 'walker.pedestrian.0002', + 'walker.pedestrian.0003', + 'walker.pedestrian.0004' + ], + 'speed_limit': 30.0 }, 'sensors': { 'vehicle_cameras': 4, @@ -406,16 +439,19 @@ def _get_default_config() -> Dict[str, Any]: 'latency_mean': 0.05, 'latency_std': 0.01, 'packet_loss_rate': 0.01, - 'message_types': ['bsm', 'spat', 'map', 'rsm', 'perception', 'warning'], + 'message_types': ['bsm', 'spat', 'map', 'rsm', 'perception', 'warning', 'pedestrian_warning'], 'update_interval': 2.0, 'security_enabled': False, 'encryption_level': 'none', - 'qos_policy': 'best_effort' + 'qos_policy': 'best_effort', + 'enable_safety_warnings': True, + 'pedestrian_warning_threshold': 10.0 }, 'cooperative': { 'num_coop_vehicles': 2, 'enable_shared_perception': True, 'enable_traffic_warnings': True, + 'enable_pedestrian_warnings': True, 'enable_maneuver_coordination': False, 'data_fusion_interval': 1.0, 'max_shared_objects': 50, @@ -431,9 +467,10 @@ def _get_default_config() -> Dict[str, Any]: 'save_enhanced': True, 'calibration_generation': True, 'enhanced_dir_name': 'enhanced', - 'methods': ['normalize', 'contrast', 'brightness'], + 'methods': ['normalize', 'contrast', 'brightness', 'pedestrian_highlight', 'safety_warning'], 'weather_effects': True, - 'augmentation_level': 'medium' + 'augmentation_level': 'medium', + 'pedestrian_safety_mode': True }, 'performance': { 'batch_size': 5, @@ -468,6 +505,7 @@ def _get_default_config() -> Dict[str, Any]: }, 'sensor_cleanup_timeout': 0.5, 'frame_rate_limit': 5.0, + 'safety_monitoring_interval': 1.0, 'memory_management': { 'gc_interval': 50, 'max_memory_mb': 500, @@ -479,16 +517,18 @@ def _get_default_config() -> Dict[str, Any]: 'output_format': 'standard', 'save_raw': True, 'save_stitched': True, - 'save_annotations': False, + 'save_annotations': True, 'save_lidar': True, 'save_fusion': True, 'save_cooperative': True, 'save_v2x_messages': True, 'save_enhanced': True, + 'save_safety_reports': True, 'validate_data': True, - 'run_analysis': False, + 'run_analysis': True, 'run_quality_check': True, 'generate_summary': True, + 'generate_safety_summary': True, 'compression_enabled': True, 'file_naming': 'sequential', 'backup_original': False @@ -501,7 +541,9 @@ def _get_default_config() -> Dict[str, Any]: 'performance_log_interval': 10.0, 'enable_progress_bar': True, 'enable_real_time_stats': True, - 'stats_update_interval': 5.0 + 'stats_update_interval': 5.0, + 'enable_safety_monitor': True, + 'safety_log_interval': 2.0 }, 'debug': { 'enable_debug_mode': False, @@ -514,7 +556,7 @@ def _get_default_config() -> Dict[str, Any]: 'metadata': { 'version': '1.0.0', 'author': 'CVIPS System', - 'description': '多传感器协同数据采集配置', + 'description': '行人安全增强数据采集配置', 'created': '', 'modified': '' } @@ -686,6 +728,7 @@ def print_config_summary(config: Dict[str, Any]): print(f" 主车: {traffic['ego_vehicles']}") print(f" 背景车辆: {traffic['background_vehicles']}") print(f" 行人: {traffic['pedestrians']}") + print(f" 车速限制: {traffic.get('speed_limit', '无')} km/h") print(f" 交通灯: {'启用' if traffic['traffic_lights'] else '禁用'}") sensors = config['sensors'] @@ -702,12 +745,13 @@ def print_config_summary(config: Dict[str, Any]): if v2x['enabled']: print(f" 通信范围: {v2x['communication_range']}米") print(f" 更新间隔: {v2x['update_interval']}秒") + print(f" 安全警告: {'启用' if v2x.get('enable_safety_warnings', False) else '禁用'}") coop = config['cooperative'] print(f"\n🤝 协同感知:") print(f" 协同车辆: {coop['num_coop_vehicles']}") print(f" 共享感知: {'启用' if coop['enable_shared_perception'] else '禁用'}") - print(f" 交通警告: {'启用' if coop['enable_traffic_warnings'] else '禁用'}") + print(f" 行人警告: {'启用' if coop.get('enable_pedestrian_warnings', False) else '禁用'}") perf = config['performance'] print(f"\n⚡ 性能:") @@ -715,6 +759,7 @@ def print_config_summary(config: Dict[str, Any]): print(f" 压缩: {'启用' if perf['enable_compression'] else '禁用'}") print(f" 下采样: {'启用' if perf['enable_downsampling'] else '禁用'}") print(f" 帧率限制: {perf['frame_rate_limit']} FPS") + print(f" 安全监控间隔: {perf.get('safety_monitoring_interval', 1.0)}秒") output = config['output'] print(f"\n💾 输出:") @@ -724,6 +769,10 @@ def print_config_summary(config: Dict[str, Any]): if isinstance(v, bool) and v and k.startswith('save_')] print(f" 启用输出: {', '.join(enabled_outputs)}") + print(f"\n🛡️ 行人安全:") + print(f" 安全监控: {'启用' if config['monitoring'].get('enable_safety_monitor', False) else '禁用'}") + print(f" 增强安全模式: {'启用' if config['enhancement'].get('pedestrian_safety_mode', False) else '禁用'}") + print("=" * 60) @staticmethod diff --git a/src/enhance_pedestrian_safety/main.py b/src/enhance_pedestrian_safety/main.py index 2626fe0a82..bfd920d8f4 100644 --- a/src/enhance_pedestrian_safety/main.py +++ b/src/enhance_pedestrian_safety/main.py @@ -1,4 +1,4 @@ -#!/usr/bin/env python3 +# !/usr/bin/env python3 import sys import os import time @@ -17,7 +17,6 @@ from config_manager import ConfigManager from annotation_generator import AnnotationGenerator from data_validator import DataValidator -from data_analyzer import DataAnalyzer from lidar_processor import LidarProcessor, MultiSensorFusion from multi_vehicle_manager import MultiVehicleManager from pedestrian_safety_monitor import PedestrianSafetyMonitor @@ -488,7 +487,7 @@ def _spawn_vehicles(self): def _spawn_pedestrians(self, center_location): blueprint_lib = self.world.get_blueprint_library() - num_peds = min(self.config['traffic']['pedestrians'], 8) + num_peds = min(self.config['traffic']['pedestrians'], 12) # 增加行人数量 spawned = 0 for _ in range(num_peds): @@ -499,8 +498,13 @@ def _spawn_pedestrians(self, center_location): ped_bp = random.choice(ped_bps) - angle = random.uniform(0, 2 * math.pi) - distance = random.uniform(5.0, 12.0) + # 在学校区域或人行横道场景中,行人更集中 + if self.config.get('scenario', {}).get('name', '').lower() in ['school_zone', 'pedestrian_crossing']: + angle = random.uniform(0, 2 * math.pi) + distance = random.uniform(3.0, 8.0) # 更靠近中心 + else: + angle = random.uniform(0, 2 * math.pi) + distance = random.uniform(5.0, 15.0) location = carla.Location( x=center_location.x + distance * math.cos(angle), @@ -1090,7 +1094,6 @@ def setup_scene(self): {'type': 'vehicle', 'capabilities': ['bsm', 'rsm']} ) - # 初始化行人安全监控器 self.safety_monitor = PedestrianSafetyMonitor(self.world, self.output_dir) time.sleep(3.0) @@ -1183,7 +1186,9 @@ def collect_data(self): if current_time - last_safety_check >= 1.0 and self.safety_monitor: safety_report = self.safety_monitor.check_pedestrian_safety() if safety_report['risk_distribution']['high'] > 0: - Log.safety(f"行人安全警告: {safety_report['high_risk_cases']}个高风险情况") + Log.safety(f"行人安全警告: {safety_report['risk_distribution']['high']}个高风险情况") + # 广播行人警告 + self._broadcast_pedestrian_warnings(safety_report) last_safety_check = current_time if current_time - last_performance_sample >= 10.0: @@ -1253,6 +1258,7 @@ def collect_data(self): final_report = self.safety_monitor.generate_final_report() Log.safety( f"行人安全报告: {final_report['risk_distribution']['high']}高风险, {final_report['risk_distribution']['medium']}中风险") + Log.safety(f"行人安全评分: {final_report['safety_score']:.1f}/100") else: Log.warning("未收集到任何数据帧") @@ -1262,6 +1268,44 @@ def collect_data(self): if self.output_format != 'standard': self._convert_to_target_format() + def _broadcast_pedestrian_warnings(self, safety_report): + """广播行人警告""" + if not self.multi_vehicle_manager: + return + + # 检查高风险交互 + if safety_report.get('risk_distribution', {}).get('high', 0) > 0: + for vehicle in self.ego_vehicles + self.multi_vehicle_manager.cooperative_vehicles: + if not hasattr(vehicle, 'is_alive') or not vehicle.is_alive: + continue + + try: + location = vehicle.get_location() + velocity = vehicle.get_velocity() + speed = math.sqrt(velocity.x ** 2 + velocity.y ** 2 + velocity.z ** 2) + + # 模拟行人位置(实际应用中应从感知系统获取) + pedestrian_location = ( + location.x + random.uniform(-5, 5), + location.y + random.uniform(-5, 5), + location.z + ) + + distance = math.sqrt( + (location.x - pedestrian_location[0]) ** 2 + + (location.y - pedestrian_location[1]) ** 2 + ) + + if distance < 20.0: # 只广播近距离行人 + self.multi_vehicle_manager.share_pedestrian_warning( + vehicle.id, + pedestrian_location, + distance, + speed + ) + except: + pass + def _force_memory_cleanup(self): Log.info("执行强制内存清理...") @@ -1646,17 +1690,12 @@ def run_validation(self): Log.info("运行数据验证...") DataValidator.validate_dataset(self.output_dir) - def run_analysis(self): - if self.config['output'].get('run_analysis', False) and self.collected_frames > 0: - Log.info("运行数据分析...") - DataAnalyzer.analyze_dataset(self.output_dir) - def main(): - parser = argparse.ArgumentParser(description='CVIPS 性能优化数据收集系统') + parser = argparse.ArgumentParser(description='CVIPS 行人安全增强数据收集系统') parser.add_argument('--config', type=str, help='配置文件路径') - parser.add_argument('--scenario', type=str, default='performance_optimized', help='场景名称') + parser.add_argument('--scenario', type=str, default='pedestrian_safety', help='场景名称') parser.add_argument('--town', type=str, default='Town10HD', choices=['Town03', 'Town04', 'Town05', 'Town10HD'], help='地图') parser.add_argument('--weather', type=str, default='clear', @@ -1665,7 +1704,7 @@ def main(): choices=['noon', 'sunset', 'night'], help='时间') parser.add_argument('--num-vehicles', type=int, default=8, help='背景车辆数') - parser.add_argument('--num-pedestrians', type=int, default=6, help='行人数') + parser.add_argument('--num-pedestrians', type=int, default=12, help='行人数') parser.add_argument('--num-coop-vehicles', type=int, default=2, help='协同车辆数') parser.add_argument('--duration', type=int, default=60, help='收集时长(秒)') @@ -1702,7 +1741,7 @@ def main(): config['output']['output_format'] = args.output_format print("\n" + "=" * 60) - print("CVIPS 性能优化数据收集系统 - 行人安全增强版") + print("CVIPS 行人安全增强数据收集系统") print("=" * 60) print(f"场景: {config['scenario']['name']}") @@ -1747,7 +1786,6 @@ def main(): collector.collect_data() - collector.run_analysis() collector.run_validation() except KeyboardInterrupt: diff --git a/src/enhance_pedestrian_safety/multi_vehicle_manager.py b/src/enhance_pedestrian_safety/multi_vehicle_manager.py index 62fcd37cf8..bb9edaafb5 100644 --- a/src/enhance_pedestrian_safety/multi_vehicle_manager.py +++ b/src/enhance_pedestrian_safety/multi_vehicle_manager.py @@ -40,7 +40,7 @@ def __init__(self, world, config, output_dir): self.config = config self.output_dir = output_dir - # 创建协同数据目录 - 确保所有子目录都创建 + # 创建协同数据目录 self.coop_dir = os.path.join(output_dir, "cooperative") os.makedirs(self.coop_dir, exist_ok=True) @@ -59,8 +59,8 @@ def __init__(self, world, config, output_dir): # V2X通信 self.v2x_messages = [] - self.communication_range = config.get('v2x', {}).get('communication_range', 300.0) # 通信范围 - self.message_buffer = defaultdict(list) # 车辆ID -> 接收的消息列表 + self.communication_range = config.get('v2x', {}).get('communication_range', 300.0) + self.message_buffer = defaultdict(list) # 协同感知 self.shared_objects = [] # 共享的感知对象 @@ -71,7 +71,9 @@ def __init__(self, world, config, output_dir): 'total_messages': 0, 'successful_transmissions': 0, 'collaborative_detections': 0, - 'data_exchange_mb': 0.0 + 'data_exchange_mb': 0.0, + 'pedestrian_warnings_sent': 0, + 'safety_alerts': 0 } def spawn_cooperative_vehicles(self, num_vehicles: int = 3) -> List[carla.Actor]: @@ -83,7 +85,6 @@ def spawn_cooperative_vehicles(self, num_vehicles: int = 3) -> List[carla.Actor] print("警告:无生成点") return [] - # 车辆类型 vehicle_types = [ 'vehicle.tesla.model3', 'vehicle.audi.tt', @@ -96,17 +97,13 @@ def spawn_cooperative_vehicles(self, num_vehicles: int = 3) -> List[carla.Actor] for i in range(min(num_vehicles, len(spawn_points))): try: - # 选择车辆类型 vtype = random.choice(vehicle_types) vehicle_bp = random.choice(blueprint_lib.filter(vtype)) - # 设置车辆属性 vehicle_bp.set_attribute('role_name', f'coop_vehicle_{i}') - # 选择生成点 spawn_point = spawn_points[i % len(spawn_points)] - # 调整位置,使车辆不在同一位置 offset_x = random.uniform(-5.0, 5.0) offset_y = random.uniform(-5.0, 5.0) location = carla.Location( @@ -115,7 +112,6 @@ def spawn_cooperative_vehicles(self, num_vehicles: int = 3) -> List[carla.Actor] z=spawn_point.location.z ) - # 轻微调整朝向 rotation = carla.Rotation( pitch=spawn_point.rotation.pitch, yaw=spawn_point.rotation.yaw + random.uniform(-15, 15), @@ -124,17 +120,12 @@ def spawn_cooperative_vehicles(self, num_vehicles: int = 3) -> List[carla.Actor] transform = carla.Transform(location, rotation) - # 生成车辆 vehicle = self.world.spawn_actor(vehicle_bp, transform) - - # 设置自动驾驶 vehicle.set_autopilot(True) - # 添加到列表 self.cooperative_vehicles.append(vehicle) spawned_vehicles.append(vehicle) - # 初始化车辆状态 self.vehicle_states[vehicle.id] = VehicleState( vehicle_id=vehicle.id, type_id=vehicle.type_id, @@ -183,7 +174,7 @@ def create_v2x_message(self, sender_id: int, message_type: str, data: Dict, message_type=message_type, data=data, timestamp=time.time(), - ttl=5.0, # 5秒生存时间 + ttl=5.0, priority=priority ) @@ -202,13 +193,11 @@ def broadcast_message(self, message: V2XMessage): recipients = [] - # 检查所有车辆是否在通信范围内 for vehicle in self.ego_vehicles + self.cooperative_vehicles: if vehicle.id != message.sender_id and vehicle.id in self.vehicle_states: receiver_state = self.vehicle_states[vehicle.id] receiver_pos = receiver_state.position - # 计算距离 distance = math.sqrt( (sender_pos[0] - receiver_pos[0]) ** 2 + (sender_pos[1] - receiver_pos[1]) ** 2 + @@ -218,7 +207,6 @@ def broadcast_message(self, message: V2XMessage): if distance <= self.communication_range: recipients.append(vehicle.id) - # 添加到接收者消息缓冲区 self.message_buffer[vehicle.id].append({ 'message': message, 'receive_time': time.time(), @@ -228,7 +216,6 @@ def broadcast_message(self, message: V2XMessage): if recipients: self.stats['successful_transmissions'] += len(recipients) - # 保存消息 try: self._save_v2x_message(message, recipients) except Exception as e: @@ -242,7 +229,6 @@ def share_perception_data(self, vehicle_id: int, detected_objects: List[Dict]): return None try: - # 创建感知消息 perception_data = { 'vehicle_id': vehicle_id, 'timestamp': time.time(), @@ -255,13 +241,11 @@ def share_perception_data(self, vehicle_id: int, detected_objects: List[Dict]): vehicle_id, 'perception', perception_data, - priority=2 # 感知数据优先级较高 + priority=2 ) - # 广播消息 recipients = self.broadcast_message(message) - # 融合共享的感知数据 if recipients: self._fuse_shared_perception(vehicle_id, detected_objects, recipients) @@ -276,7 +260,7 @@ def share_traffic_warning(self, vehicle_id: int, warning_type: str, """共享交通警告""" try: warning_data = { - 'warning_type': warning_type, # 'accident', 'congestion', 'hazard', 'construction' + 'warning_type': warning_type, 'location': location, 'severity': severity, 'timestamp': time.time(), @@ -287,7 +271,7 @@ def share_traffic_warning(self, vehicle_id: int, warning_type: str, vehicle_id, 'warning', warning_data, - priority=3 # 警告消息优先级最高 + priority=3 ) self.broadcast_message(message) @@ -298,9 +282,19 @@ def share_traffic_warning(self, vehicle_id: int, warning_type: str, return None def share_pedestrian_warning(self, vehicle_id: int, pedestrian_location: Tuple[float, float, float], - distance: float, speed: float): + distance: float, speed: float, pedestrian_id: Optional[int] = None): """共享行人警告""" try: + # 风险评估 + if distance < 5.0: + severity = 'critical' + elif distance < 10.0: + severity = 'high' + elif distance < 20.0: + severity = 'medium' + else: + severity = 'low' + warning_data = { 'warning_type': 'pedestrian', 'pedestrian_location': pedestrian_location, @@ -308,7 +302,9 @@ def share_pedestrian_warning(self, vehicle_id: int, pedestrian_location: Tuple[f 'vehicle_speed': speed, 'timestamp': time.time(), 'source_vehicle': vehicle_id, - 'severity': 'high' if distance < 10.0 else 'medium' if distance < 20.0 else 'low' + 'pedestrian_id': pedestrian_id, + 'severity': severity, + 'recommended_action': self._get_recommended_action(distance, speed, severity) } message = self.create_v2x_message( @@ -320,27 +316,39 @@ def share_pedestrian_warning(self, vehicle_id: int, pedestrian_location: Tuple[f recipients = self.broadcast_message(message) + # 更新统计 + self.stats['pedestrian_warnings_sent'] += 1 + if severity in ['critical', 'high']: + self.stats['safety_alerts'] += 1 + return message, recipients except Exception as e: print(f"共享行人警告失败: {e}") return None, [] + def _get_recommended_action(self, distance: float, speed: float, severity: str) -> str: + """获取推荐的安全措施""" + if severity == 'critical': + return "立即紧急制动,准备避让" + elif severity == 'high': + return "减速至20km/h以下,准备制动" + elif severity == 'medium': + return "减速至30km/h,保持警惕" + else: + return "保持当前速度,注意观察" + def _fuse_shared_perception(self, source_id: int, objects: List[Dict], recipients: List[int]): """融合共享的感知数据""" fused_objects = [] for obj in objects: - # 转换为全局坐标系(简化处理,实际需要坐标变换) global_obj = obj.copy() global_obj['source_vehicles'] = [source_id] - global_obj['confidence'] = obj.get('confidence', 0.8) # 降低置信度 + global_obj['confidence'] = obj.get('confidence', 0.8) - # 检查是否已有类似对象 matched = False for existing_obj in self.shared_objects: - # 简单的对象匹配(基于位置) if self._objects_match(global_obj, existing_obj): - # 更新现有对象 existing_obj['source_vehicles'].append(source_id) existing_obj['confidence'] = min(1.0, existing_obj.get('confidence', 0) + 0.1) existing_obj['update_time'] = time.time() @@ -353,11 +361,10 @@ def _fuse_shared_perception(self, source_id: int, objects: List[Dict], recipient self.shared_objects.extend(fused_objects) - # 清理旧对象 current_time = time.time() self.shared_objects = [ obj for obj in self.shared_objects - if current_time - obj.get('detection_time', 0) < 10.0 # 保留10秒内的对象 + if current_time - obj.get('detection_time', 0) < 10.0 ] if fused_objects: @@ -368,7 +375,6 @@ def _objects_match(self, obj1: Dict, obj2: Dict, distance_threshold: float = 5.0 if obj1.get('class') != obj2.get('class'): return False - # 计算位置距离 pos1 = obj1.get('position', {'x': 0, 'y': 0, 'z': 0}) pos2 = obj2.get('position', {'x': 0, 'y': 0, 'z': 0}) @@ -385,23 +391,20 @@ def get_shared_perception_for_vehicle(self, vehicle_id: int) -> List[Dict]: shared_data = [] for obj in self.shared_objects: - # 检查对象是否在车辆视野内(简化) shared_data.append(obj) return shared_data def coordinate_maneuvers(self, maneuvers: List[Dict]): """协调多车辆机动""" - # 分配优先级和时序 coordinated = [] for i, maneuver in enumerate(maneuvers): coordinated_maneuver = maneuver.copy() coordinated_maneuver['sequence'] = i - coordinated_maneuver['start_time'] = time.time() + i * 2.0 # 间隔2秒 + coordinated_maneuver['start_time'] = time.time() + i * 2.0 coordinated.append(coordinated_maneuver) - # 发送协调消息 message = self.create_v2x_message( maneuver.get('vehicle_id', 0), 'coordination', @@ -421,7 +424,6 @@ def _save_v2x_message(self, message: V2XMessage, recipients: List[int]): 'transmission_time': time.time() } - # 确保目录存在 v2x_messages_dir = os.path.join(self.coop_dir, "v2x_messages") os.makedirs(v2x_messages_dir, exist_ok=True) @@ -446,7 +448,6 @@ def save_shared_perception(self, frame_num: int): 'stats': self.stats } - # 确保目录存在 shared_perception_dir = os.path.join(self.coop_dir, "shared_perception") os.makedirs(shared_perception_dir, exist_ok=True) @@ -467,7 +468,12 @@ def generate_summary(self): 'shared_objects_count': len(self.shared_objects), 'communication_range': self.communication_range, 'average_messages_per_vehicle': self.stats['total_messages'] / max(1, len(self.ego_vehicles) + len( - self.cooperative_vehicles)) + self.cooperative_vehicles)), + 'safety_metrics': { + 'pedestrian_warnings': self.stats['pedestrian_warnings_sent'], + 'safety_alerts': self.stats['safety_alerts'], + 'collaborative_detections': self.stats['collaborative_detections'] + } } filepath = os.path.join(self.coop_dir, "cooperative_summary.json") diff --git a/src/enhance_pedestrian_safety/pedestrian_safety_monitor.py b/src/enhance_pedestrian_safety/pedestrian_safety_monitor.py index 5ec545f100..c5513fa9c4 100644 --- a/src/enhance_pedestrian_safety/pedestrian_safety_monitor.py +++ b/src/enhance_pedestrian_safety/pedestrian_safety_monitor.py @@ -19,17 +19,22 @@ def __init__(self, world, output_dir): # 安全参数 self.safety_thresholds = { + 'critical_distance': 2.0, # 临界距离 (米) 'high_risk_distance': 5.0, # 高风险距离 (米) 'medium_risk_distance': 10.0, # 中风险距离 (米) 'low_risk_distance': 20.0, # 低风险距离 (米) 'safe_speed_limit': 30.0, # 安全速度限制 (km/h) 'reaction_time': 1.5, # 反应时间 (秒) - 'braking_deceleration': 6.0 # 制动减速度 (m/s²) + 'braking_deceleration': 6.0, # 制动减速度 (m/s²) + 'ttc_critical': 1.0, # 临界碰撞时间 (秒) + 'ttc_high': 2.0, # 高风险碰撞时间 (秒) + 'ttc_medium': 3.0 # 中风险碰撞时间 (秒) } # 统计数据 self.stats = { 'total_interactions': 0, + 'critical_cases': 0, 'high_risk_cases': 0, 'medium_risk_cases': 0, 'low_risk_cases': 0, @@ -39,12 +44,17 @@ def __init__(self, world, output_dir): 'average_distance': 0, 'min_distance': float('inf'), 'max_distance': 0, - 'interaction_times': [] + 'average_ttc': 0, + 'min_ttc': float('inf'), + 'interaction_times': [], + 'vehicle_speeds': [], + 'pedestrian_speeds': [] } # 详细记录 self.interaction_records = [] self.warning_logs = [] + self.critical_events = [] def check_pedestrian_safety(self) -> Dict: """检查行人安全""" @@ -60,19 +70,19 @@ def check_pedestrian_safety(self) -> Dict: for pedestrian in pedestrians: pedestrian_location = pedestrian.get_location() + pedestrian_velocity = pedestrian.get_velocity() # 计算距离 distance = vehicle_location.distance(pedestrian_location) # 计算相对速度 - pedestrian_velocity = pedestrian.get_velocity() relative_speed = self._calculate_relative_speed(vehicle_velocity, pedestrian_velocity) # 计算碰撞时间 time_to_collision = self._calculate_ttc(distance, relative_speed) # 评估风险 - risk_level = self._assess_risk(distance, vehicle_speed, time_to_collision) + risk_level = self._assess_risk(distance, vehicle_speed, time_to_collision, relative_speed) interaction = { 'timestamp': time.time(), @@ -80,6 +90,7 @@ def check_pedestrian_safety(self) -> Dict: 'pedestrian_id': pedestrian.id, 'distance': distance, 'vehicle_speed': vehicle_speed * 3.6, # 转换为km/h + 'pedestrian_speed': math.sqrt(pedestrian_velocity.x ** 2 + pedestrian_velocity.y ** 2) * 3.6, 'relative_speed': relative_speed * 3.6, 'time_to_collision': time_to_collision if time_to_collision < 100 else None, 'risk_level': risk_level, @@ -92,7 +103,9 @@ def check_pedestrian_safety(self) -> Dict: 'x': pedestrian_location.x, 'y': pedestrian_location.y, 'z': pedestrian_location.z - } + }, + 'safety_measures': self._suggest_safety_measures(distance, vehicle_speed, time_to_collision, + risk_level) } current_interactions.append(interaction) @@ -100,9 +113,9 @@ def check_pedestrian_safety(self) -> Dict: # 更新统计 self._update_stats(interaction) - # 记录高风险情况 - if risk_level == 'high': - self._log_high_risk(interaction) + # 记录高风险和临界情况 + if risk_level in ['critical', 'high']: + self._log_risk_event(interaction) # 保存当前检查结果 if current_interactions: @@ -124,22 +137,33 @@ def _calculate_relative_speed(self, v1: carla.Vector3D, v2: carla.Vector3D) -> f def _calculate_ttc(self, distance: float, relative_speed: float) -> float: """计算碰撞时间 (Time to Collision)""" - if relative_speed > 0.1: # 避免除以零 + if relative_speed > 0.1: return distance / relative_speed return float('inf') - def _assess_risk(self, distance: float, speed: float, ttc: Optional[float]) -> str: + def _assess_risk(self, distance: float, speed: float, ttc: Optional[float], relative_speed: float) -> str: """评估风险等级""" speed_kmh = speed * 3.6 + # 基于碰撞时间的风险评估 + if ttc is not None: + if ttc < self.safety_thresholds['ttc_critical']: + return 'critical' + elif ttc < self.safety_thresholds['ttc_high']: + return 'high' + elif ttc < self.safety_thresholds['ttc_medium']: + return 'medium' + # 基于距离的风险评估 - if distance < self.safety_thresholds['high_risk_distance']: - if ttc is not None and ttc < 2.0: + if distance < self.safety_thresholds['critical_distance']: + return 'critical' + elif distance < self.safety_thresholds['high_risk_distance']: + if speed_kmh > self.safety_thresholds['safe_speed_limit']: return 'high' else: return 'medium' elif distance < self.safety_thresholds['medium_risk_distance']: - if speed_kmh > self.safety_thresholds['safe_speed_limit']: + if speed_kmh > self.safety_thresholds['safe_speed_limit'] * 1.5: return 'medium' else: return 'low' @@ -148,10 +172,65 @@ def _assess_risk(self, distance: float, speed: float, ttc: Optional[float]) -> s else: return 'safe' + def _suggest_safety_measures(self, distance: float, speed: float, ttc: Optional[float], risk_level: str) -> List[ + str]: + """建议安全措施""" + measures = [] + speed_kmh = speed * 3.6 + + if risk_level == 'critical': + measures.extend([ + "立即紧急制动", + "鸣喇叭警告行人", + "准备紧急避让", + "向其他车辆发送紧急警告", + "记录事故数据" + ]) + elif risk_level == 'high': + measures.extend([ + "立即减速至20km/h以下", + "保持警惕,准备制动", + "观察行人动向", + "准备避让", + "向附近车辆发送警告" + ]) + elif risk_level == 'medium': + measures.extend([ + "减速至安全速度", + "保持安全距离", + "观察周围环境", + "准备应对突发情况", + "评估避让路径" + ]) + elif risk_level == 'low': + measures.extend([ + "保持当前速度", + "注意观察", + "准备减速", + "保持安全车距" + ]) + else: + measures.append("正常行驶,保持警惕") + + # 添加基于距离的额外建议 + if distance < 3.0: + measures.append("保持极端警惕,准备紧急操作") + elif distance < 5.0: + measures.append("准备随时制动") + + # 添加基于车速的额外建议 + if speed_kmh > 50.0: + measures.append("车速过快,建议减速") + elif speed_kmh > 30.0: + measures.append("注意控制车速") + + return measures + def _update_stats(self, interaction: Dict): """更新统计数据""" self.stats['total_interactions'] += 1 distance = interaction['distance'] + ttc = interaction.get('time_to_collision') # 更新距离统计 self.stats['average_distance'] = ( @@ -161,9 +240,25 @@ def _update_stats(self, interaction: Dict): self.stats['min_distance'] = min(self.stats['min_distance'], distance) self.stats['max_distance'] = max(self.stats['max_distance'], distance) + # 更新碰撞时间统计 + if ttc is not None and ttc < 100: + self.stats['average_ttc'] = ( + (self.stats['average_ttc'] * (self.stats['total_interactions'] - 1) + ttc) / + self.stats['total_interactions'] + ) + self.stats['min_ttc'] = min(self.stats['min_ttc'], ttc) + + # 更新速度统计 + self.stats['vehicle_speeds'].append(interaction['vehicle_speed']) + self.stats['pedestrian_speeds'].append(interaction['pedestrian_speed']) + # 更新风险统计 risk_level = interaction['risk_level'] - if risk_level == 'high': + if risk_level == 'critical': + self.stats['critical_cases'] += 1 + self.stats['near_misses'] += 1 + self.stats['safety_warnings'] += 1 + elif risk_level == 'high': self.stats['high_risk_cases'] += 1 self.stats['near_misses'] += 1 self.stats['safety_warnings'] += 1 @@ -185,45 +280,44 @@ def _update_stats(self, interaction: Dict): if len(self.interaction_records) > 1000: self.interaction_records = self.interaction_records[-1000:] - def _log_high_risk(self, interaction: Dict): - """记录高风险情况""" - warning = { + def _log_risk_event(self, interaction: Dict): + """记录风险事件""" + event_type = 'critical' if interaction['risk_level'] == 'critical' else 'high_risk' + + event = { 'timestamp': datetime.now().isoformat(), + 'event_type': event_type, 'interaction': interaction, - 'safety_measures': self._suggest_safety_measures(interaction) + 'safety_measures': interaction.get('safety_measures', []), + 'environment': { + 'weather': str(self.world.get_weather()), + 'time_of_day': self._get_time_of_day() + } } - self.warning_logs.append(warning) - # 保存高风险警告 + if event_type == 'critical': + self.critical_events.append(event) + # 保存临界事件 + if len(self.critical_events) % 5 == 0: + self._save_critical_events() + else: + self.warning_logs.append(event) + + # 保存警告日志 if len(self.warning_logs) % 10 == 0: self._save_warning_logs() - def _suggest_safety_measures(self, interaction: Dict) -> List[str]: - """建议安全措施""" - measures = [] - - if interaction['risk_level'] == 'high': - measures.extend([ - "立即制动", - "鸣喇叭警告", - "准备紧急避让", - "向其他车辆发送警告" - ]) - elif interaction['risk_level'] == 'medium': - measures.extend([ - "减速行驶", - "保持警惕", - "准备制动", - "观察行人动向" - ]) - elif interaction['risk_level'] == 'low': - measures.extend([ - "保持安全距离", - "观察周围环境", - "准备应对突发情况" - ]) + def _get_time_of_day(self) -> str: + """获取当前时间""" + weather = self.world.get_weather() + sun_altitude = weather.sun_altitude_angle - return measures + if sun_altitude > 45: + return 'day' + elif sun_altitude > 0: + return 'sunset' + else: + return 'night' def _save_interaction_report(self, interactions: List[Dict]): """保存交互报告""" @@ -235,10 +329,15 @@ def _save_interaction_report(self, interactions: List[Dict]): 'total_interactions': len(interactions), 'interactions': interactions, 'summary': { + 'critical': len([i for i in interactions if i['risk_level'] == 'critical']), 'high_risk': len([i for i in interactions if i['risk_level'] == 'high']), 'medium_risk': len([i for i in interactions if i['risk_level'] == 'medium']), 'low_risk': len([i for i in interactions if i['risk_level'] == 'low']), 'safe': len([i for i in interactions if i['risk_level'] == 'safe']) + }, + 'environment': { + 'weather': str(self.world.get_weather()), + 'time_of_day': self._get_time_of_day() } } @@ -254,6 +353,15 @@ def _save_warning_logs(self): with open(warning_file, 'w', encoding='utf-8') as f: json.dump(self.warning_logs, f, indent=2, ensure_ascii=False) + def _save_critical_events(self): + """保存临界事件""" + if not self.critical_events: + return + + critical_file = os.path.join(self.safety_dir, "critical_events.json") + with open(critical_file, 'w', encoding='utf-8') as f: + json.dump(self.critical_events, f, indent=2, ensure_ascii=False) + def _generate_safety_report(self) -> Dict: """生成安全报告""" report = { @@ -261,13 +369,22 @@ def _generate_safety_report(self) -> Dict: 'statistics': self.stats.copy(), 'safety_thresholds': self.safety_thresholds, 'risk_distribution': { + 'critical': self.stats['critical_cases'], 'high': self.stats['high_risk_cases'], 'medium': self.stats['medium_risk_cases'], 'low': self.stats['low_risk_cases'], 'safe': self.stats['safe_cases'] }, 'safety_score': self._calculate_safety_score(), - 'recommendations': self._generate_recommendations() + 'recommendations': self._generate_recommendations(), + 'performance_metrics': { + 'average_vehicle_speed': np.mean(self.stats['vehicle_speeds']) if self.stats['vehicle_speeds'] else 0, + 'average_pedestrian_speed': np.mean(self.stats['pedestrian_speeds']) if self.stats[ + 'pedestrian_speeds'] else 0, + 'max_vehicle_speed': max(self.stats['vehicle_speeds']) if self.stats['vehicle_speeds'] else 0, + 'interaction_frequency': len(self.stats['interaction_times']) / max(1, ( + time.time() - min(self.stats['interaction_times']) if self.stats['interaction_times'] else 1)) + } } return report @@ -277,17 +394,30 @@ def _calculate_safety_score(self) -> float: if self.stats['total_interactions'] == 0: return 100.0 + critical_ratio = self.stats['critical_cases'] / self.stats['total_interactions'] high_risk_ratio = self.stats['high_risk_cases'] / self.stats['total_interactions'] medium_risk_ratio = self.stats['medium_risk_cases'] / self.stats['total_interactions'] - # 评分公式:基础分减去风险比例 - score = 100 - (high_risk_ratio * 60 + medium_risk_ratio * 30) * 100 + # 评分公式:基础分减去风险比例加权 + score = 100 - (critical_ratio * 80 + high_risk_ratio * 50 + medium_risk_ratio * 20) * 100 # 考虑平均距离 if self.stats['average_distance'] > 15.0: score += 10 elif self.stats['average_distance'] < 5.0: score -= 20 + elif self.stats['average_distance'] < 2.0: + score -= 40 + + # 考虑平均车速 + if self.stats['vehicle_speeds']: + avg_speed = np.mean(self.stats['vehicle_speeds']) + if avg_speed > 50.0: + score -= 15 + elif avg_speed > 30.0: + score -= 5 + elif avg_speed < 20.0: + score += 10 return max(0, min(100, score)) @@ -295,19 +425,37 @@ def _generate_recommendations(self) -> List[str]: """生成改进建议""" recommendations = [] - if self.stats['high_risk_cases'] > 0: + if self.stats['critical_cases'] > 0: + recommendations.extend([ + "立即改进行人检测系统", + "增加紧急制动系统", + "实施更严格的限速措施", + "增加行人预警系统", + "记录和分析所有临界事件" + ]) + + if self.stats['high_risk_cases'] > 5: recommendations.extend([ "增加行人安全距离阈值", "加强车辆行人检测系统", - "实施更严格的限速措施", - "增加行人警告系统" + "实施自动减速系统", + "增加行人警告频率" ]) if self.stats['average_distance'] < 10.0: recommendations.append("增加车辆与行人的平均距离") - if self.stats['near_misses'] > 5: - recommendations.append("实施紧急制动系统") + if self.stats['near_misses'] > 3: + recommendations.append("实施紧急避让系统") + + if self.stats['vehicle_speeds']: + avg_speed = np.mean(self.stats['vehicle_speeds']) + if avg_speed > 40.0: + recommendations.append("降低平均车速,特别是在行人密集区域") + + # 基于碰撞时间的建议 + if self.stats['min_ttc'] < 2.0: + recommendations.append("改进碰撞时间预测算法") return recommendations @@ -319,9 +467,32 @@ def generate_final_report(self) -> Dict: final_report['historical_data'] = { 'total_interaction_records': len(self.interaction_records), 'total_warning_logs': len(self.warning_logs), + 'total_critical_events': len(self.critical_events), 'analysis_period': self._get_analysis_period() } + # 添加详细分析 + final_report['detailed_analysis'] = { + 'distance_analysis': { + 'average': self.stats['average_distance'], + 'min': self.stats['min_distance'], + 'max': self.stats['max_distance'], + 'std': np.std([r['distance'] for r in self.interaction_records]) if self.interaction_records else 0 + }, + 'speed_analysis': { + 'vehicle_avg': np.mean(self.stats['vehicle_speeds']) if self.stats['vehicle_speeds'] else 0, + 'pedestrian_avg': np.mean(self.stats['pedestrian_speeds']) if self.stats['pedestrian_speeds'] else 0, + 'vehicle_std': np.std(self.stats['vehicle_speeds']) if self.stats['vehicle_speeds'] else 0 + }, + 'ttc_analysis': { + 'average': self.stats['average_ttc'], + 'min': self.stats['min_ttc'], + 'percentage_below_2s': len( + [r for r in self.interaction_records if r.get('time_to_collision', 100) < 2.0]) / max(1, + len(self.interaction_records)) * 100 + } + } + # 保存最终报告 final_file = os.path.join(self.safety_dir, "final_safety_report.json") with open(final_file, 'w', encoding='utf-8') as f: @@ -359,10 +530,14 @@ def save_data(self): with open(records_file, 'w', encoding='utf-8') as f: json.dump(self.interaction_records, f, indent=2, ensure_ascii=False) - # 保存警告日志 + # 保存警告日志和临界事件 self._save_warning_logs() + self._save_critical_events() # 生成并保存最终报告 - self.generate_final_report() + final_report = self.generate_final_report() - print(f"行人安全数据已保存到: {self.safety_dir}") \ No newline at end of file + print(f"行人安全数据已保存到: {self.safety_dir}") + print(f"安全评分: {final_report['safety_score']:.1f}/100") + print(f"高风险事件: {final_report['risk_distribution']['high']}次") + print(f"临界事件: {final_report['risk_distribution']['critical']}次") \ No newline at end of file diff --git a/src/enhance_pedestrian_safety/scene_manager.py b/src/enhance_pedestrian_safety/scene_manager.py index f7e081f9ed..f43601c4fd 100644 --- a/src/enhance_pedestrian_safety/scene_manager.py +++ b/src/enhance_pedestrian_safety/scene_manager.py @@ -52,8 +52,8 @@ class SceneManager: 'vehicle_density': 0.4, 'pedestrian_density': 0.9, 'weather_variations': ['clear', 'cloudy'], - 'speed_limit': 20.0, - 'safety_zone_radius': 30.0 + 'speed_limit': 15.0, # 降低限速 + 'safety_zone_radius': 50.0 # 扩大安全区域 }, 'pedestrian_crossing': { 'description': '人行横道场景', @@ -94,14 +94,16 @@ def setup_scene(world, config, scene_type='intersection_4way'): # 如果是学校区域,设置车速限制 if scene_type == 'school_zone' and 'speed_limit' in scene_config: config['traffic']['speed_limit'] = scene_config['speed_limit'] + # 添加行人安全区域标记 + SceneManager._add_safety_zone_markers(world, config) # 应用场景特定设置 - SceneManager._apply_scene_specifics(world, scene_type) + SceneManager._apply_scene_specifics(world, scene_type, config) return config @staticmethod - def _apply_scene_specifics(world, scene_type): + def _apply_scene_specifics(world, scene_type, config): """应用场景特定的设置""" try: if scene_type == 'highway': @@ -134,6 +136,8 @@ def _apply_scene_specifics(world, scene_type): actor.set_light_state(carla.VehicleLightState.LowBeam) except: pass + # 添加学校区域警告标志 + SceneManager._add_school_zone_signs(world) elif scene_type == 'pedestrian_crossing': # 人行横道:增加可见性 @@ -143,10 +147,71 @@ def _apply_scene_specifics(world, scene_type): actor.set_light_state(carla.VehicleLightState.LowBeam) except: pass + # 添加人行横道标记 + SceneManager._add_crosswalk_markings(world) except Exception as e: print(f"场景特定设置失败: {e}") + @staticmethod + def _add_safety_zone_markers(world, config): + """添加安全区域标记""" + try: + blueprint_lib = world.get_blueprint_library() + + # 在学校区域周围添加安全锥 + spawn_points = world.get_map().get_spawn_points() + if spawn_points: + center_point = spawn_points[len(spawn_points) // 2].location + SceneManager.spawn_traffic_cones(world, center_point, num_cones=15) + + # 添加警告标志 + warning_sign_bp = blueprint_lib.find('static.prop.trafficsign') + if warning_sign_bp: + sign_location = carla.Location(center_point.x, center_point.y, center_point.z + 2.0) + world.spawn_actor(warning_sign_bp, carla.Transform(sign_location, carla.Rotation(0, 90, 0))) + + except Exception as e: + print(f"添加安全区域标记失败: {e}") + + @staticmethod + def _add_school_zone_signs(world): + """添加学校区域标志""" + try: + blueprint_lib = world.get_blueprint_library() + + # 查找学校区域标志蓝图 + school_sign_bp = None + for bp in blueprint_lib.filter('static.prop.*'): + if 'school' in bp.id.lower() or 'warning' in bp.id.lower(): + school_sign_bp = bp + break + + if school_sign_bp: + # 在学校区域周围放置标志 + spawn_points = world.get_map().get_spawn_points() + if spawn_points: + for i in range(min(4, len(spawn_points))): + location = spawn_points[i].location + sign_location = carla.Location(location.x, location.y, location.z + 2.0) + world.spawn_actor(school_sign_bp, + carla.Transform(sign_location, carla.Rotation(0, i * 90, 0))) + + except Exception as e: + print(f"添加学校区域标志失败: {e}") + + @staticmethod + def _add_crosswalk_markings(world): + """添加人行横道标记""" + try: + spawn_points = world.get_map().get_spawn_points() + if spawn_points: + center_point = spawn_points[len(spawn_points) // 2].location + SceneManager.spawn_pedestrian_safety_features(world, center_point, 'crosswalk') + + except Exception as e: + print(f"添加人行横道标记失败: {e}") + @staticmethod def spawn_pedestrian_safety_features(world, location, feature_type='crosswalk'): """生成行人安全设施""" @@ -315,7 +380,12 @@ def save_scene_description(output_dir, scene_type, config, extra_info=None): 'description': SceneManager.SCENES.get(scene_type, {}).get('description', '未知场景'), 'config': config, 'created': SceneManager._get_timestamp(), - 'extra_info': extra_info or {} + 'extra_info': extra_info or {}, + 'safety_features': { + 'has_crosswalk': scene_type in ['school_zone', 'pedestrian_crossing', 'intersection_4way'], + 'has_traffic_cones': scene_type in ['school_zone', 'construction_zone'], + 'has_warning_signs': scene_type in ['school_zone', 'pedestrian_crossing'] + } } scene_file = os.path.join(output_dir, "metadata", "scene_description.json") diff --git a/src/enhance_pedestrian_safety/sensor_enhancer.py b/src/enhance_pedestrian_safety/sensor_enhancer.py index 3dba1d767b..4f2d2b8f27 100644 --- a/src/enhance_pedestrian_safety/sensor_enhancer.py +++ b/src/enhance_pedestrian_safety/sensor_enhancer.py @@ -1,5 +1,5 @@ """ -传感器数据增强模块 - 提高数据质量和多样性(优化版) +传感器数据增强模块 - 提高数据质量和多样性(行人安全优化版) """ import numpy as np @@ -7,7 +7,6 @@ import random import os import json -from PIL import Image, ImageEnhance, ImageFilter from datetime import datetime from typing import Dict, List, Tuple, Optional, Union, Callable import concurrent.futures @@ -46,6 +45,8 @@ class EnhancementMethod(Enum): COLOR_TEMP = "color_temperature" JPEG_COMPRESSION = "jpeg_compression" COLOR_JITTER = "color_jitter" + PEDESTRIAN_HIGHLIGHT = "pedestrian_highlight" # 新增:行人高亮 + SAFETY_WARNING = "safety_warning" # 新增:安全警告 @dataclass @@ -61,6 +62,7 @@ class EnhancementConfig: save_enhanced: bool = True output_format: str = "jpg" compression_quality: int = 90 + pedestrian_safety_mode: bool = True # 新增:行人安全模式 def __post_init__(self): if self.enabled_methods is None: @@ -69,6 +71,12 @@ def __post_init__(self): EnhancementMethod.CONTRAST, EnhancementMethod.BRIGHTNESS ] + # 如果启用行人安全模式,添加相关方法 + if self.pedestrian_safety_mode: + if EnhancementMethod.PEDESTRIAN_HIGHLIGHT not in self.enabled_methods: + self.enabled_methods.append(EnhancementMethod.PEDESTRIAN_HIGHLIGHT) + if EnhancementMethod.SAFETY_WARNING not in self.enabled_methods: + self.enabled_methods.append(EnhancementMethod.SAFETY_WARNING) class BatchEnhancer: @@ -193,7 +201,7 @@ def get_stats(self) -> Dict: class SensorDataEnhancer: - """传感器数据增强器(优化版)""" + """传感器数据增强器(行人安全优化版)""" def __init__(self, config: Union[EnhancementConfig, Dict]): if isinstance(config, dict): @@ -210,9 +218,13 @@ def __init__(self, config: Union[EnhancementConfig, Dict]): 'method_times': {} } + # 行人检测模拟数据 + self.pedestrian_detections = [] + self.safety_warnings = [] + def _setup_method_registry(self) -> Dict[EnhancementMethod, Callable]: """设置方法注册表""" - return { + registry = { EnhancementMethod.NORMALIZE: self._normalize_image, EnhancementMethod.BRIGHTNESS: self._adjust_brightness, EnhancementMethod.CONTRAST: self._adjust_contrast, @@ -228,8 +240,11 @@ def _setup_method_registry(self) -> Dict[EnhancementMethod, Callable]: EnhancementMethod.VIGNETTE: self._add_vignette, EnhancementMethod.COLOR_TEMP: self._adjust_color_temperature, EnhancementMethod.JPEG_COMPRESSION: self._simulate_jpeg_compression, - EnhancementMethod.COLOR_JITTER: self._color_jitter + EnhancementMethod.COLOR_JITTER: self._color_jitter, + EnhancementMethod.PEDESTRIAN_HIGHLIGHT: self._highlight_pedestrians, # 新增 + EnhancementMethod.SAFETY_WARNING: self._add_safety_warnings # 新增 } + return registry def _setup_weather_methods(self) -> Dict[WeatherType, List[EnhancementMethod]]: """设置天气相关方法""" @@ -262,7 +277,8 @@ def _setup_weather_methods(self) -> Dict[WeatherType, List[EnhancementMethod]]: EnhancementMethod.BRIGHTNESS, EnhancementMethod.CONTRAST, EnhancementMethod.NOISE, - EnhancementMethod.VIGNETTE + EnhancementMethod.VIGNETTE, + EnhancementMethod.PEDESTRIAN_HIGHLIGHT # 夜间特别关注行人 ], WeatherType.SUNSET: [ EnhancementMethod.NORMALIZE, @@ -276,7 +292,7 @@ def enhance_image(self, image_data: np.ndarray, sensor_type: str = 'camera', return_methods: bool = False) -> Union[np.ndarray, Tuple]: """ - 增强图像数据(优化版) + 增强图像数据(行人安全优化版) Args: image_data: 原始图像数据 (H, W, C) @@ -404,11 +420,10 @@ def get_performance_stats(self) -> Dict: stats['avg_time_per_call'] = stats['total_time'] / stats['calls'] return stats - # ========== 增强方法实现(优化版)========== + # ========== 增强方法实现(行人安全优化版)========== def _normalize_image(self, image: np.ndarray) -> np.ndarray: """图像归一化(优化版)""" - # 使用OpenCV加速 if image.dtype != np.uint8: normalized = cv2.normalize(image, None, 0, 255, cv2.NORM_MINMAX) return normalized.astype(np.uint8) @@ -418,9 +433,7 @@ def _adjust_brightness(self, image: np.ndarray) -> np.ndarray: """调整亮度(优化版)""" factor = random.uniform(*self.config.intensity_range) - # 使用NumPy向量化操作 if factor != 1.0: - # 转换为浮点数进行计算 img_float = image.astype(np.float32) * factor return np.clip(img_float, 0, 255).astype(np.uint8) return image @@ -430,9 +443,7 @@ def _adjust_contrast(self, image: np.ndarray) -> np.ndarray: factor = random.uniform(*self.config.intensity_range) if factor != 1.0: - # 计算平均值 mean = np.mean(image, axis=(0, 1), keepdims=True) - # 应用对比度调整 contrasted = mean + factor * (image.astype(np.float32) - mean) return np.clip(contrasted, 0, 255).astype(np.uint8) return image @@ -442,18 +453,14 @@ def _adjust_saturation(self, image: np.ndarray) -> np.ndarray: factor = random.uniform(*self.config.intensity_range) if factor != 1.0: - # 转换为HSV空间 hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV).astype(np.float32) - # 调整饱和度通道 hsv[:, :, 1] = np.clip(hsv[:, :, 1] * factor, 0, 255) - # 转换回RGB saturated = cv2.cvtColor(hsv.astype(np.uint8), cv2.COLOR_HSV2RGB) return saturated return image def _enhance_sharpness(self, image: np.ndarray) -> np.ndarray: """增强锐度(优化版)""" - # 使用拉普拉斯算子增强边缘 kernel = np.array([[0, -1, 0], [-1, 5, -1], [0, -1, 0]]) @@ -464,7 +471,6 @@ def _gamma_correction(self, image: np.ndarray) -> np.ndarray: """伽马校正(优化版)""" gamma = random.uniform(0.5, 2.0) - # 使用LUT加速 table = np.array([((i / 255.0) ** gamma) * 255 for i in range(256)]).astype(np.uint8) corrected = cv2.LUT(image, table) return corrected @@ -474,7 +480,6 @@ def _add_noise(self, image: np.ndarray) -> np.ndarray: noise_type = random.choice(['gaussian', 'salt_pepper']) if noise_type == 'gaussian': - # 高斯噪声 mean = 0 var = random.uniform(0.001, 0.005) sigma = var ** 0.5 @@ -483,18 +488,15 @@ def _add_noise(self, image: np.ndarray) -> np.ndarray: return np.clip(noisy, 0, 255).astype(np.uint8) else: # salt_pepper - # 椒盐噪声 amount = random.uniform(0.001, 0.005) s_vs_p = random.uniform(0.3, 0.7) noisy = image.copy() - # 盐噪声 num_salt = np.ceil(amount * image.size * s_vs_p / 3) coords = [np.random.randint(0, i, int(num_salt)) for i in image.shape] noisy[coords[0], coords[1], :] = 255 - # 椒噪声 num_pepper = np.ceil(amount * image.size * (1.0 - s_vs_p) / 3) coords = [np.random.randint(0, i, int(num_pepper)) for i in image.shape] noisy[coords[0], coords[1], :] = 0 @@ -513,7 +515,6 @@ def _apply_motion_blur(self, image: np.ndarray) -> np.ndarray: kernel_size = random.choice([7, 9, 11]) direction = random.choice(['horizontal', 'vertical', 'diagonal']) - # 创建运动模糊核 kernel = np.zeros((kernel_size, kernel_size)) if direction == 'horizontal': @@ -526,7 +527,6 @@ def _apply_motion_blur(self, image: np.ndarray) -> np.ndarray: blurred = cv2.filter2D(image, -1, kernel) - # 混合原图和模糊图 alpha = random.uniform(0.3, 0.7) result = cv2.addWeighted(image, 1 - alpha, blurred, alpha, 0) return result.astype(np.uint8) @@ -535,10 +535,8 @@ def _add_rain_effect(self, image: np.ndarray) -> np.ndarray: """添加雨滴效果(优化版)""" h, w = image.shape[:2] - # 创建雨滴层 rain_layer = np.zeros((h, w), dtype=np.float32) - # 生成随机雨滴 num_drops = random.randint(200, 800) drop_length = random.randint(8, 15) @@ -547,15 +545,12 @@ def _add_rain_effect(self, image: np.ndarray) -> np.ndarray: y = random.randint(0, h - 1) brightness = random.uniform(0.3, 0.6) - # 绘制雨滴线 for i in range(drop_length): if y + i < h and x + i < w: rain_layer[y + i, x + i] += brightness - # 模糊雨滴层 rain_layer = cv2.GaussianBlur(rain_layer, (3, 3), 0) - # 叠加到原图 rain_layer_3d = np.stack([rain_layer] * 3, axis=2) enhanced = image.astype(np.float32) * 0.9 + rain_layer_3d * 0.1 * 255 return np.clip(enhanced, 0, 255).astype(np.uint8) @@ -564,17 +559,14 @@ def _add_fog_effect(self, image: np.ndarray) -> np.ndarray: """添加雾效(优化版)""" h, w = image.shape[:2] - # 创建深度图(简化版,假设中心最近) center_y, center_x = h // 2, w // 2 y_coords, x_coords = np.ogrid[:h, :w] distance = np.sqrt((x_coords - center_x) ** 2 + (y_coords - center_y) ** 2) distance = distance / np.max(distance) - # 雾效强度 fog_intensity = random.uniform(0.2, 0.5) fog_color = random.choice([200, 210, 220]) - # 应用雾效 fog_strength = distance * fog_intensity fog_strength_3d = np.stack([fog_strength] * 3, axis=2) fog_layer = np.ones_like(image) * fog_color @@ -586,19 +578,15 @@ def _add_fog_effect(self, image: np.ndarray) -> np.ndarray: def _add_cloud_effect(self, image: np.ndarray) -> np.ndarray: """添加云层效果(优化版)""" - # 降低饱和度和对比度,模拟多云天气 hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV).astype(np.float32) - # 调整饱和度 hsv[:, :, 1] *= random.uniform(0.7, 0.9) - # 调整亮度和对比度 brightness_factor = random.uniform(0.9, 1.1) contrast_factor = random.uniform(0.8, 0.95) hsv[:, :, 2] = np.clip(hsv[:, :, 2] * brightness_factor, 0, 255) - # 应用对比度调整 mean = np.mean(hsv[:, :, 2]) hsv[:, :, 2] = mean + contrast_factor * (hsv[:, :, 2] - mean) hsv[:, :, 2] = np.clip(hsv[:, :, 2], 0, 255) @@ -610,21 +598,17 @@ def _add_vignette(self, image: np.ndarray) -> np.ndarray: """添加暗角效果(优化版)""" h, w = image.shape[:2] - # 创建暗角蒙版 center_y, center_x = h // 2, w // 2 y_coords, x_coords = np.ogrid[:h, :w] - # 计算距离(使用椭圆形状) y_dist = (y_coords - center_y) / (h / 2) x_dist = (x_coords - center_x) / (w / 2) distance = np.sqrt(x_dist ** 2 + y_dist ** 2) - # 创建暗角 vignette_intensity = random.uniform(0.2, 0.4) vignette = 1 - distance * vignette_intensity vignette = np.clip(vignette, 0.6, 1.0) - # 应用暗角 vignette_3d = np.stack([vignette] * 3, axis=2) enhanced = image.astype(np.float32) * vignette_3d @@ -634,16 +618,13 @@ def _adjust_color_temperature(self, image: np.ndarray) -> np.ndarray: """调整色温(优化版)""" temp_type = random.choice(['warm', 'cool']) - # 转换为浮点数 img_float = image.astype(np.float32) if temp_type == 'warm': - # 暖色调:增加红色和黄色 img_float[:, :, 0] *= random.uniform(1.0, 1.15) # 红色 img_float[:, :, 1] *= random.uniform(1.0, 1.1) # 绿色 img_float[:, :, 2] *= random.uniform(0.9, 1.0) # 蓝色 else: - # 冷色调:增加蓝色 img_float[:, :, 0] *= random.uniform(0.9, 1.0) # 红色 img_float[:, :, 1] *= random.uniform(0.95, 1.0) # 绿色 img_float[:, :, 2] *= random.uniform(1.0, 1.15) # 蓝色 @@ -653,15 +634,12 @@ def _adjust_color_temperature(self, image: np.ndarray) -> np.ndarray: def _simulate_jpeg_compression(self, image: np.ndarray) -> np.ndarray: """模拟JPEG压缩(优化版)""" - # 使用OpenCV的JPEG编码/解码模拟压缩 quality = random.randint(70, 95) - # 编码为JPEG encode_param = [int(cv2.IMWRITE_JPEG_QUALITY), quality] result, encoded = cv2.imencode('.jpg', cv2.cvtColor(image, cv2.COLOR_RGB2BGR), encode_param) if result: - # 解码 decoded = cv2.imdecode(encoded, cv2.IMREAD_COLOR) return cv2.cvtColor(decoded, cv2.COLOR_BGR2RGB) @@ -669,57 +647,124 @@ def _simulate_jpeg_compression(self, image: np.ndarray) -> np.ndarray: def _color_jitter(self, image: np.ndarray) -> np.ndarray: """颜色抖动(优化版)""" - # 随机应用多种颜色变换 transforms = [] - # 亮度调整 if random.random() > 0.5: brightness = random.uniform(0.8, 1.2) transforms.append(lambda img: self._adjust_brightness(img, brightness)) - # 对比度调整 if random.random() > 0.5: contrast = random.uniform(0.8, 1.2) transforms.append(lambda img: self._adjust_contrast(img, contrast)) - # 饱和度调整 if random.random() > 0.5: saturation = random.uniform(0.8, 1.2) transforms.append(lambda img: self._adjust_saturation(img, saturation)) - # 随机应用变换 if transforms: random.shuffle(transforms) - for transform in transforms[:2]: # 最多应用2个变换 + for transform in transforms[:2]: image = transform(image) return image + def _highlight_pedestrians(self, image: np.ndarray) -> np.ndarray: + """高亮行人(新增方法)""" + if not self.config.pedestrian_safety_mode: + return image + + h, w = image.shape[:2] + highlighted = image.copy() + + # 模拟行人检测(随机位置) + num_pedestrians = random.randint(1, 3) + + for i in range(num_pedestrians): + # 随机位置 + x = random.randint(50, w - 150) + y = random.randint(50, h - 200) + + # 随机大小 + width = random.randint(30, 70) + height = random.randint(80, 180) + + # 根据距离设置颜色(模拟风险评估) + distance = random.uniform(5.0, 50.0) + if distance < 10.0: + color = (0, 0, 255) # 红色:高风险 + thickness = 3 + elif distance < 20.0: + color = (0, 165, 255) # 橙色:中风险 + thickness = 2 + else: + color = (0, 255, 0) # 绿色:低风险 + thickness = 1 + + # 绘制边界框 + cv2.rectangle(highlighted, (x, y), (x + width, y + height), color, thickness) + + # 添加标签 + label = f"Pedestrian {distance:.1f}m" + cv2.putText(highlighted, label, (x, y - 10), + cv2.FONT_HERSHEY_SIMPLEX, 0.5, color, thickness) + + # 记录检测 + self.pedestrian_detections.append({ + 'position': (x, y, width, height), + 'distance': distance, + 'risk_level': 'high' if distance < 10.0 else 'medium' if distance < 20.0 else 'low' + }) + + return highlighted + + def _add_safety_warnings(self, image: np.ndarray) -> np.ndarray: + """添加安全警告(新增方法)""" + if not self.config.pedestrian_safety_mode: + return image + + warned = image.copy() + + # 检查是否有高风险行人 + high_risk = [d for d in self.pedestrian_detections if d['risk_level'] == 'high'] + + if high_risk: + # 在图像顶部添加警告条 + warning_height = 40 + warning_bar = np.zeros((warning_height, image.shape[1], 3), dtype=np.uint8) + warning_bar[:] = (0, 0, 255) # 红色背景 + + # 添加警告文本 + warning_text = f"WARNING: {len(high_risk)} high-risk pedestrians detected!" + cv2.putText(warning_bar, warning_text, (10, 25), + cv2.FONT_HERSHEY_SIMPLEX, 0.7, (255, 255, 255), 2) + + # 将警告条添加到图像顶部 + warned = np.vstack([warning_bar, warned]) + + # 记录警告 + self.safety_warnings.append({ + 'timestamp': time.time(), + 'high_risk_count': len(high_risk), + 'message': warning_text + }) + + return warned + def _enhance_non_rgb_image(self, image_data: np.ndarray, sensor_type: str) -> np.ndarray: """增强非RGB图像(深度、语义分割)""" if sensor_type == 'depth': - # 深度图增强:归一化和去噪 if image_data.dtype != np.uint8: - # 归一化到0-255 normalized = cv2.normalize(image_data, None, 0, 255, cv2.NORM_MINMAX) else: normalized = image_data - # 应用中值滤波去除噪声 enhanced = cv2.medianBlur(normalized, 3) return enhanced.astype(np.uint8) elif sensor_type == 'semantic': - # 语义分割图增强:保持类别不变,只做边界平滑 kernel = np.ones((3, 3), np.uint8) - - # 形态学操作:先腐蚀再膨胀(闭运算)填充小孔 enhanced = cv2.morphologyEx(image_data, cv2.MORPH_CLOSE, kernel) - - # 高斯模糊平滑边界 enhanced = cv2.GaussianBlur(enhanced, (3, 3), 0.5) - - # 恢复类别值 enhanced = np.round(enhanced).astype(image_data.dtype) return enhanced @@ -727,19 +772,15 @@ def _enhance_non_rgb_image(self, image_data: np.ndarray, sensor_type: str) -> np def save_enhanced_image(self, image_data: np.ndarray, output_path: str, metadata: Optional[Dict] = None): - """保存增强后的图像和元数据(优化版)""" + """保存增强后的图像和元数据""" try: - # 确保目录存在 os.makedirs(os.path.dirname(output_path), exist_ok=True) - # 保存图像 if output_path.lower().endswith(('.png', '.jpg', '.jpeg')): cv2.imwrite(output_path, cv2.cvtColor(image_data, cv2.COLOR_RGB2BGR)) else: - # 默认保存为PNG cv2.imwrite(output_path + '.png', cv2.cvtColor(image_data, cv2.COLOR_RGB2BGR)) - # 保存元数据 if metadata: meta_path = output_path.rsplit('.', 1)[0] + '_meta.json' with open(meta_path, 'w', encoding='utf-8') as f: @@ -750,7 +791,7 @@ def save_enhanced_image(self, image_data: np.ndarray, output_path: str, raise def generate_enhancement_report(self, output_dir: str) -> Dict: - """生成增强报告(优化版)""" + """生成增强报告""" report = { 'timestamp': datetime.now().isoformat(), 'config': { @@ -758,7 +799,8 @@ def generate_enhancement_report(self, output_dir: str) -> Dict: 'time_of_day': self.config.time_of_day, 'enabled_methods': [m.value for m in self.config.enabled_methods], 'intensity_range': self.config.intensity_range, - 'probability': self.config.probability + 'probability': self.config.probability, + 'pedestrian_safety_mode': self.config.pedestrian_safety_mode }, 'performance_stats': self.get_performance_stats(), 'method_usage': { @@ -768,153 +810,20 @@ def generate_enhancement_report(self, output_dir: str) -> Dict: 'cache_info': { 'cache_size': len(self.method_cache), 'cache_hits': self.perf_stats.get('cache_hit_calls', 0) + }, + 'safety_statistics': { + 'pedestrian_detections': len(self.pedestrian_detections), + 'safety_warnings': len(self.safety_warnings), + 'high_risk_cases': len([d for d in self.pedestrian_detections if d['risk_level'] == 'high']) } } - # 保存报告 report_path = os.path.join(output_dir, 'enhancement_report.json') with open(report_path, 'w', encoding='utf-8') as f: json.dump(report, f, indent=2, ensure_ascii=False) return report - def _add_pedestrian_detection_markers(self, image: np.ndarray, pedestrians: List[Dict]) -> np.ndarray: - """添加行人检测标记""" - if not pedestrians: - return image - - # 创建副本 - marked_image = image.copy() - - for pedestrian in pedestrians: - # 模拟行人检测框 - x = random.randint(100, image.shape[1] - 100) - y = random.randint(100, image.shape[0] - 100) - width = random.randint(30, 80) - height = random.randint(80, 180) - - # 绘制检测框 - color = (0, 255, 0) # 绿色表示安全 - thickness = 2 - - cv2.rectangle(marked_image, (x, y), (x + width, y + height), color, thickness) - - # 添加标签 - label = f"Pedestrian {pedestrian.get('distance', 0):.1f}m" - cv2.putText(marked_image, label, (x, y - 10), - cv2.FONT_HERSHEY_SIMPLEX, 0.5, color, thickness) - - return marked_image - - def _simulate_safety_warnings(self, image: np.ndarray, warnings: List[Dict]) -> np.ndarray: - """模拟安全警告""" - if not warnings: - return image - - warning_image = image.copy() - - for warning in warnings[:2]: # 最多显示2个警告 - # 添加警告文本 - text = f"WARNING: {warning.get('type', 'Pedestrian')} {warning.get('distance', 0):.1f}m" - color = (0, 0, 255) # 红色表示警告 - thickness = 2 - - position = (50, 50 + warnings.index(warning) * 40) - cv2.putText(warning_image, text, position, - cv2.FONT_HERSHEY_SIMPLEX, 0.7, color, thickness) - - # 添加警告图标 - icon_size = 30 - icon_position = (position[0] - icon_size - 10, position[1] - icon_size // 2) - cv2.rectangle(warning_image, icon_position, - (icon_position[0] + icon_size, icon_position[1] + icon_size), - color, -1) # 填充红色矩形 - - return warning_image - - -# 保持向后兼容的类 -class SensorCalibrator: - """传感器校准模块(优化版)""" - - def __init__(self, config: Dict): - self.config = config - self.calibration_data = {} - - def generate_calibration_files(self, output_dir: str, - vehicle_locations: List[Dict], - camera_positions: List[Dict]): - """生成传感器校准文件(优化版)""" - calib_dir = os.path.join(output_dir, "calibration") - os.makedirs(calib_dir, exist_ok=True) - - # 1. 生成相机内参 - self._generate_camera_intrinsics(calib_dir) - - # 2. 生成外参(相机到车辆) - self._generate_extrinsics(calib_dir, vehicle_locations, camera_positions) - - # 3. 生成传感器间标定 - self._generate_sensor_calibration(calib_dir) - - # 4. 生成时间同步校准 - self._generate_temporal_calibration(calib_dir) - - # 5. 生成验证数据 - self._generate_validation_data(calib_dir) - - print(f"校准文件已生成到: {calib_dir}") - return calib_dir - - def _generate_camera_intrinsics(self, calib_dir: str): - """生成相机内参(优化版)""" - image_size = self.config.get('sensors', {}).get('image_size', [1280, 720]) - width, height = image_size[0], image_size[1] - - # 为不同相机生成不同的内参 - camera_types = [ - ('front_wide', 100.0), # 前视广角 - ('front_narrow', 60.0), # 前视窄角 - ('side', 90.0), # 侧视 - ('rear', 120.0), # 后视 - ('infrastructure', 120.0) # 基础设施 - ] - - for camera_name, fov in camera_types: - # 内参矩阵 [fx, 0, cx; 0, fy, cy; 0, 0, 1] - fx = width / (2 * np.tan(np.radians(fov / 2))) - fy = fx - cx = width / 2.0 - cy = height / 2.0 - - intrinsics = { - 'camera_name': camera_name, - 'camera_matrix': [ - [float(fx), 0.0, float(cx)], - [0.0, float(fy), float(cy)], - [0.0, 0.0, 1.0] - ], - 'distortion_coefficients': [ - random.uniform(-0.1, 0.1), # k1 - random.uniform(-0.01, 0.01), # k2 - 0.0, # p1 - 0.0, # p2 - random.uniform(-0.001, 0.001) # k3 - ], - 'image_size': [int(width), int(height)], - 'fov': float(fov), - 'pixel_size': [0.003, 0.003], # 假设像素大小 - 'sensor_type': 'pinhole', - 'calibration_date': datetime.now().isoformat(), - 'accuracy': random.uniform(0.5, 1.0) # 标定精度(像素) - } - - file_path = os.path.join(calib_dir, f'{camera_name}_intrinsic.json') - with open(file_path, 'w', encoding='utf-8') as f: - json.dump(intrinsics, f, indent=2, ensure_ascii=False) - - # ... 其他方法保持不变,但可以添加更多优化 ... - # 兼容旧版本接口 def enhance_image(image_data, sensor_type='camera'): diff --git a/src/enhance_pedestrian_safety/v2x_communication.py b/src/enhance_pedestrian_safety/v2x_communication.py index 63079cc46a..937c00d15e 100644 --- a/src/enhance_pedestrian_safety/v2x_communication.py +++ b/src/enhance_pedestrian_safety/v2x_communication.py @@ -29,32 +29,29 @@ def __init__(self, config: Dict): self.config = config self.enabled = config.get('enabled', True) - # 通信参数 - self.communication_range = config.get('communication_range', 300.0) # 通信范围(米) - self.bandwidth = config.get('bandwidth', 10.0) # 带宽(Mbps) - self.latency_mean = config.get('latency_mean', 0.05) # 平均延迟(秒) - self.latency_std = config.get('latency_std', 0.01) # 延迟标准差 - self.packet_loss_rate = config.get('packet_loss_rate', 0.01) # 丢包率 - - # 消息管理 + self.communication_range = config.get('communication_range', 300.0) + self.bandwidth = config.get('bandwidth', 10.0) + self.latency_mean = config.get('latency_mean', 0.05) + self.latency_std = config.get('latency_std', 0.01) + self.packet_loss_rate = config.get('packet_loss_rate', 0.01) + self.message_queue = queue.PriorityQueue() self.received_messages = [] self.message_counter = 0 - # 网络模拟 - self.network_nodes = {} # 节点ID -> 节点信息 - self.connections = {} # 连接状态 + self.network_nodes = {} + self.connections = {} - # 统计信息 self.stats = { 'messages_sent': 0, 'messages_received': 0, 'messages_dropped': 0, 'total_latency': 0.0, - 'bandwidth_used': 0.0 + 'bandwidth_used': 0.0, + 'pedestrian_warnings': 0, + 'safety_alerts': 0 } - # 启动消息处理线程 if self.enabled: self.running = True self.processor_thread = threading.Thread(target=self._message_processor, daemon=True) @@ -79,17 +76,14 @@ def send_message(self, message: V2XMessage) -> bool: if not self.enabled: return False - # 模拟网络延迟和丢包 if np.random.random() < self.packet_loss_rate: self.stats['messages_dropped'] += 1 return False - # 添加延迟 latency = np.random.normal(self.latency_mean, self.latency_std) delivery_time = time.time() + max(0, latency) - # 添加到消息队列(按优先级和发送时间排序) - priority = -message.priority # 负号因为PriorityQueue是小顶堆 + priority = -message.priority self.message_queue.put((priority, delivery_time, message)) self.stats['messages_sent'] += 1 @@ -162,7 +156,7 @@ def broadcast_map_data(self, map_id: str, map_data: Dict) -> Optional[V2XMessage message_type='map', data=map_data, timestamp=time.time(), - ttl=30.0, # 地图数据TTL较长 + ttl=30.0, priority=1, position=map_data.get('reference_point') ) @@ -183,13 +177,18 @@ def broadcast_roadside_safety_message(self, rsu_id: str, warning_data: Dict) -> data=warning_data, timestamp=time.time(), ttl=3.0, - priority=3, # 安全消息优先级高 + priority=3, position=warning_data.get('event_position') ) self.message_counter += 1 if self.send_message(message): + # 更新安全统计 + if warning_data.get('type') == 'pedestrian': + self.stats['pedestrian_warnings'] += 1 + if warning_data.get('severity') in ['high', 'critical']: + self.stats['safety_alerts'] += 1 return message return None @@ -202,13 +201,11 @@ def get_reachable_nodes(self, sender_id: str, sender_position: tuple) -> List[st if node_id == sender_id: continue - # 计算距离 distance = self._calculate_distance(sender_position, node_info['position']) if distance <= self.communication_range: - # 模拟信号衰减 signal_strength = 1.0 - (distance / self.communication_range) - if signal_strength > 0.3: # 最小信号强度阈值 + if signal_strength > 0.3: reachable.append(node_id) return reachable @@ -228,19 +225,16 @@ def _message_processor(self): """消息处理线程""" while self.running: try: - # 获取下一个要传递的消息 if not self.message_queue.empty(): priority, delivery_time, message = self.message_queue.get_nowait() current_time = time.time() if current_time >= delivery_time: - # 消息已到传递时间 self._deliver_message(message) else: - # 重新放回队列 self.message_queue.put((priority, delivery_time, message)) - time.sleep(0.001) # 短暂休眠 + time.sleep(0.001) else: time.sleep(0.01) @@ -254,18 +248,15 @@ def _deliver_message(self, message: V2XMessage): if not message.position: return - # 获取可达节点 reachable_nodes = self.get_reachable_nodes(message.sender_id, message.position) - # 模拟带宽限制 - message_size = len(json.dumps(asdict(message)).encode('utf-8')) / 1024 / 1024 # MB - bandwidth_required = message_size * 8 * len(reachable_nodes) # Mbps + message_size = len(json.dumps(asdict(message)).encode('utf-8')) / 1024 / 1024 + bandwidth_required = message_size * 8 * len(reachable_nodes) - if bandwidth_required > self.bandwidth * 0.8: # 带宽超过80%时随机丢包 + if bandwidth_required > self.bandwidth * 0.8: drop_probability = min(1.0, bandwidth_required / self.bandwidth - 0.8) reachable_nodes = [n for n in reachable_nodes if np.random.random() > drop_probability] - # 添加到接收消息列表 for node_id in reachable_nodes: received_msg = { 'message': asdict(message), @@ -278,7 +269,6 @@ def _deliver_message(self, message: V2XMessage): self.received_messages.append(received_msg) self.stats['messages_received'] += 1 - # 更新带宽使用统计 self.stats['bandwidth_used'] += bandwidth_required def get_messages_for_node(self, node_id: str, message_types: List[str] = None) -> List[Dict]: @@ -291,7 +281,6 @@ def get_messages_for_node(self, node_id: str, message_types: List[str] = None) - if not message_types or message['message_type'] in message_types: messages.append(msg_record) - # 清理已处理的消息 self.received_messages = [ msg for msg in self.received_messages if msg['receiver_id'] != node_id or @@ -308,7 +297,11 @@ def get_network_status(self) -> Dict: 'bandwidth_utilization': (self.stats['bandwidth_used'] / ( self.bandwidth * time.time())) * 100 if time.time() > 0 else 0, 'message_delivery_rate': self.stats['messages_received'] / max(1, self.stats['messages_sent']), - 'average_latency': self.stats['total_latency'] / max(1, self.stats['messages_sent']) + 'average_latency': self.stats['total_latency'] / max(1, self.stats['messages_sent']), + 'safety_metrics': { + 'pedestrian_warnings': self.stats['pedestrian_warnings'], + 'safety_alerts': self.stats['safety_alerts'] + } } def reset_stats(self): @@ -318,7 +311,9 @@ def reset_stats(self): 'messages_received': 0, 'messages_dropped': 0, 'total_latency': 0.0, - 'bandwidth_used': 0.0 + 'bandwidth_used': 0.0, + 'pedestrian_warnings': 0, + 'safety_alerts': 0 } def stop(self):