use std::collections::HashMap;
use anyhow::{anyhow, Result};
use crate::nn::policy::Policy;
use crate::rl::env::Env;
pub struct CollectedData {
pub obs: Vec<Vec<usize>>,
pub logits: Vec<Vec<f32>>,
pub perms: Vec<Option<usize>>,
pub values: Vec<f32>,
pub rewards: Vec<f32>,
pub actions: Vec<usize>,
pub additional_data: HashMap<String, Vec<f32>>,
}
pub fn merge(mut chunks: Vec<CollectedData>) -> Result<CollectedData> {
let mut merged = chunks.pop().ok_or(anyhow!("Something went wrong. No data in collected data chunks to merge. "))?;
for chunk in chunks {
merged.merge(&chunk);
}
Ok(merged)
}
impl CollectedData {
pub fn new(
obs: Vec<Vec<usize>>,
logits: Vec<Vec<f32>>,
perms: Vec<Option<usize>>,
values: Vec<f32>,
rewards: Vec<f32>,
actions: Vec<usize>,
) -> Self {
CollectedData {
obs,
logits,
perms,
values,
rewards,
actions,
additional_data: HashMap::new(),
}
}
pub fn merge(&mut self, other: &CollectedData) {
self.obs.extend(other.obs.iter().cloned());
self.logits.extend(other.logits.iter().cloned());
self.perms.extend(other.perms.iter().cloned());
self.values.extend(&other.values);
self.rewards.extend(&other.rewards);
self.actions.extend(&other.actions);
for (key, value_vec) in &other.additional_data {
self.additional_data
.entry(key.clone())
.and_modify(|existing| existing.extend(value_vec.iter().cloned()))
.or_insert_with(|| value_vec.clone());
}
}
}
pub trait Collector: Send + Sync {
fn collect(&self, env: &Box<dyn Env>, policy: &Policy) -> Result<CollectedData>;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_merge_collected_data() {
let d1 = CollectedData::new(
vec![vec![0]],
vec![vec![0.1]],
vec![Some(0)],
vec![0.2],
vec![0.3],
vec![1],
);
let d2 = CollectedData::new(
vec![vec![1]],
vec![vec![0.4]],
vec![None],
vec![0.5],
vec![0.6],
vec![0],
);
let merged = merge(vec![d1, d2]).unwrap();
assert_eq!(merged.obs.len(), 2);
assert_eq!(merged.logits.len(), 2);
assert_eq!(merged.actions, vec![0, 1]);
}
}