Rust泛型函数convert_to_one_hot如何设置T为usize或Vec<usize>的trait约束
实现方案
Rust 不支持直接限定泛型为多个指定类型的联合,我们可以通过自定义 trait 为允许的输入类型封装统一行为,即可实现你需要的泛型约束。
另外注意你给出的示例代码存在一个可运行问题:Vec::with_capacity 只会预分配内存空间,不会初始化元素,直接通过下标访问会触发 panic,我们可以用 vec![false; num_classes] 一步完成向量的初始化和默认值赋值。
1 定义统一行为的自定义 trait
我们定义一个 trait 封装「将自身对应类别设置到独热向量」的行为:
trait ToOneHotCategory { fn set_to_one_hot(self, one_hot: &mut Vec<bool>); }
2 为允许的输入类型实现 trait
分别为 usize 和 Vec<usize> 实现上面定义的 trait:
impl ToOneHotCategory for usize { fn set_to_one_hot(self, one_hot: &mut Vec<bool>) { one_hot[self] = true; } } impl ToOneHotCategory for Vec<usize> { fn set_to_one_hot(self, one_hot: &mut Vec<bool>) { for category in self { one_hot[category] = true; } } }
3 实现泛型函数
给泛型参数 T 加上 ToOneHotCategory 约束即可:
pub fn convert_to_one_hot<T: ToOneHotCategory>(category_id: T, num_classes: usize) -> Vec<bool> { let mut one_hot = vec![false; num_classes]; category_id.set_to_one_hot(&mut one_hot); one_hot }
调用示例
fn main() { // 输入为 usize 类型 let single_one_hot = convert_to_one_hot(2, 5); assert_eq!(single_one_hot, vec![false, false, true, false, false]); // 输入为 Vec<usize> 类型 let multi_one_hot = convert_to_one_hot(vec![1, 3], 5); assert_eq!(multi_one_hot, vec![false, true, false, true, false]); }
拓展说明
如果你需要支持引用类型输入(避免转移所有权),可以额外添加引用类型的 trait 实现:
impl<'a> ToOneHotCategory for &'a usize { fn set_to_one_hot(self, one_hot: &mut Vec<bool>) { one_hot[*self] = true; } } impl<'a> ToOneHotCategory for &'a Vec<usize> { fn set_to_one_hot(self, one_hot: &mut Vec<bool>) { for category in self { one_hot[*category] = true; } } }
内容的提问来源于stack exchange,提问作者ujjwal.researcher
相关产品推荐
相关产品推荐

