use std::collections::HashMap;
use crate::autograd::Variable;
use crate::nn::{self, Buffer, Module, Parameter};
use crate::tensor::{Result, TensorError};
use super::{Graph, GraphExt};
use super::trend::Trend;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PathKind {
Subgraph,
Tag,
}
#[allow(dead_code)]
pub(crate) enum ResolvedPath<'a> {
Subgraph(&'a Graph),
Tag { graph: &'a Graph, tag: String },
}
impl Graph {
pub(crate) fn resolve(&self, path: &str) -> Result<ResolvedPath<'_>> {
if path.is_empty() {
return Err(TensorError::new("empty label path"));
}
let segments: Vec<&str> = path.split('.').collect();
self.resolve_segments(&segments, path, false)
}
fn resolve_segments<'a>(
&'a self,
segments: &[&str],
full_path: &str,
cross_boundary: bool,
) -> Result<ResolvedPath<'a>> {
debug_assert!(!segments.is_empty());
let first = segments[0];
if segments.len() == 1 {
if let Some(g) = self.child_graph(first) {
return Ok(ResolvedPath::Subgraph(g));
}
if self.tag_names.contains_key(first) {
if cross_boundary && self.internal_tags.contains(first) {
return Err(TensorError::new(&format!(
"tag {:?} is internal and cannot be accessed from a parent graph (path: {:?})",
first, full_path
)));
}
return Ok(ResolvedPath::Tag { graph: self, tag: first.to_string() });
}
return Err(TensorError::new(&format!(
"{:?} is not a subgraph or tag of this graph (path: {:?})",
first, full_path
)));
}
let child = self.child_graph(first).ok_or_else(|| {
TensorError::new(&format!(
"{:?} is not a subgraph of this graph (path: {:?})",
first, full_path
))
})?;
child.resolve_segments(&segments[1..], full_path, true)
}
pub fn tree_children(&self) -> HashMap<&str, &Graph> {
self.children.iter()
.filter_map(|(label, &ni)| {
self.nodes[ni].module.as_ref()
.and_then(|m| m.as_graph())
.map(|g| (label.as_str(), g))
})
.collect()
}
pub fn child_graph(&self, label: &str) -> Option<&Graph> {
self.children.get(label)
.and_then(|&ni| self.nodes[ni].module.as_ref())
.and_then(|m| m.as_graph())
}
pub fn subgraph(&self, path: &str) -> Result<&Graph> {
match self.resolve(path)? {
ResolvedPath::Subgraph(g) => Ok(g),
ResolvedPath::Tag { .. } => Err(TensorError::new(&format!(
"path {:?} resolves to a tag, not a subgraph", path
))),
}
}
pub fn is_composed(&self) -> bool {
self.composed.get()
}
pub fn internal_tags(&self) -> &std::collections::HashSet<String> {
&self.internal_tags
}
pub fn validate_path(&self, path: &str) -> Result<PathKind> {
match self.resolve(path)? {
ResolvedPath::Subgraph(_) => Ok(PathKind::Subgraph),
ResolvedPath::Tag { .. } => Ok(PathKind::Tag),
}
}
pub fn parameters_at(&self, path: &str) -> Result<Vec<Parameter>> {
match self.resolve(path)? {
ResolvedPath::Subgraph(g) => Ok(g.parameters()),
ResolvedPath::Tag { graph, ref tag } => {
if let Some(&(ni, _)) = graph.tag_names.get(tag.as_str()) {
if let Some(ref module) = graph.nodes[ni].module {
Ok(module.parameters())
} else {
Ok(vec![])
}
} else {
Ok(vec![])
}
}
}
}
pub fn named_parameters_at(&self, path: &str) -> Result<Vec<(String, Parameter)>> {
match self.resolve(path)? {
ResolvedPath::Subgraph(g) => Ok(g.named_parameters()),
ResolvedPath::Tag { graph, ref tag } => {
if let Some(&(ni, _)) = graph.tag_names.get(tag.as_str()) {
if let Some(ref module) = graph.nodes[ni].module {
Ok(module.parameters().into_iter()
.map(|p| (format!("{}/{}", tag, p.name), p))
.collect())
} else {
Ok(vec![])
}
} else {
Ok(vec![])
}
}
}
}
pub fn named_buffers_at(&self, path: &str) -> Result<Vec<(String, Buffer)>> {
match self.resolve(path)? {
ResolvedPath::Subgraph(g) => Ok(g.named_buffers()),
ResolvedPath::Tag { graph, ref tag } => {
if let Some(&(ni, _)) = graph.tag_names.get(tag.as_str()) {
if let Some(ref module) = graph.nodes[ni].module {
Ok(module.buffers().into_iter()
.map(|b| (format!("{}/{}", tag, b.name), b))
.collect())
} else {
Ok(vec![])
}
} else {
Ok(vec![])
}
}
}
}
pub fn freeze(&self, path: &str) -> Result<()> {
for p in self.parameters_at(path)? {
p.freeze()?;
}
Ok(())
}
pub fn thaw(&self, path: &str) -> Result<()> {
for p in self.parameters_at(path)? {
p.unfreeze()?;
}
Ok(())
}
pub fn is_frozen(&self, path: &str) -> Result<bool> {
let params = self.parameters_at(path)?;
if params.is_empty() {
return Ok(false);
}
Ok(params.iter().all(|p| p.is_frozen()))
}
pub fn load_subgraph_checkpoint(&self, path: &str, file: &str) -> Result<nn::LoadReport> {
let target = self.subgraph(path)?;
let params = target.named_parameters();
let buffers = target.named_buffers();
let hash = target.structural_hash();
nn::load_checkpoint_file(file, ¶ms, &buffers, Some(hash))
}
pub fn set_training_at(&self, path: &str, training: bool) -> Result<()> {
match self.resolve(path)? {
ResolvedPath::Subgraph(g) => {
g.set_training(training);
}
ResolvedPath::Tag { graph, ref tag } => {
if let Some(&(ni, _)) = graph.tag_names.get(tag.as_str()) {
if let Some(ref module) = graph.nodes[ni].module {
crate::nn::walk_modules(module.as_ref(), &mut |m| {
m.set_training(training);
});
}
}
}
}
Ok(())
}
pub fn tagged_at(&self, path: &str) -> Result<Option<Variable>> {
match self.resolve(path)? {
ResolvedPath::Subgraph(_) => Err(TensorError::new(&format!(
"path {:?} resolves to a subgraph, not a tag", path
))),
ResolvedPath::Tag { graph, ref tag } => Ok(graph.tagged(tag)),
}
}
pub fn collect_at(&self, paths: &[&str]) -> Result<()> {
for &path in paths {
match self.resolve(path)? {
ResolvedPath::Subgraph(_) => {
return Err(TensorError::new(&format!(
"collect_at: {:?} resolves to a subgraph, not a tag", path
)));
}
ResolvedPath::Tag { graph, ref tag } => {
graph.collect(&[tag.as_str()])?;
}
}
}
Ok(())
}
pub fn record_at(&self, path: &str, value: f64) -> Result<()> {
let segments: Vec<&str> = path.split('.').collect();
if segments.len() < 2 {
self.record_scalar(path, value);
return Ok(());
}
let parent_path = segments[..segments.len() - 1].join(".");
let tag = segments[segments.len() - 1];
let target = self.subgraph(&parent_path)?;
target.record_scalar(tag, value);
Ok(())
}
pub fn trend_at(&self, path: &str) -> Result<Trend> {
let segments: Vec<&str> = path.split('.').collect();
if segments.len() < 2 {
return Ok(self.trend(path));
}
let parent_path = segments[..segments.len() - 1].join(".");
let tag = segments[segments.len() - 1];
let target = self.subgraph(&parent_path)?;
Ok(target.trend(tag))
}
}
#[cfg(test)]
#[path = "tree_tests.rs"]
mod tests;