use std::f32;
use std::fmt;
#[derive(Debug)]
pub struct Result {
cells:Vec<String>
}
impl Result{
fn new()->Result{
Result{cells:Vec::new(),}
}
}
#[derive(Clone,Copy,Debug)]
pub enum ColType {
Conti, Discr,}
#[derive(Debug)]
pub struct Stati {
k:String, n:usize, child:Vec<Stati>,
}
impl Stati{
fn new(k:String)->Stati{
Stati{
k,
n:0,
child:Vec::new(),
}
}
}
#[derive(Debug)]
pub struct Col {
cells:Vec<String>,
name:String,
t:ColType,
stati:Stati,
min_gini:f32,
devide_val:f32,
sel_val:String,
}
impl Col{
fn new(name:String,t:ColType)->Col{
Col{
cells:Vec::new(),
name:name.clone(),
t,
stati:Stati::new(name),
min_gini:-1.0,
devide_val:-1.0,
sel_val:String::from(""),
}
}
}
#[derive(Debug)]
pub struct D {
cols:Vec<Col>,
r:Result,
}
impl D {
fn get_conti_min_gini_and_devide_val(&mut self) {
for c in self.cols.iter_mut(){
if let ColType::Conti = c.t{
let mut copy = Vec::new();
for vi in c.cells.iter() {
let f:f32 = vi.parse().unwrap();
copy.push(f);
}
copy = sort(&mut copy);
let mut mid = Vec::new();
for i in 0..(copy.len()-1) {
let c1 = copy[i];
let c2 = copy[i+1];
if c1 == c2 {
continue;
}
let m = (c1 + c2) / 2.0;
mid.push(m);
}
let mut min_gini = f32::MAX;
let mut devide_val = -1.0_f32;
for m in mid.iter() {
let mut sta = Stati::new(c.name.clone());
for i in 0..c.cells.len() {
let f:f32 = c.cells[i].parse().unwrap();
sta.n += 1;
let mut k = format!("<{}",m);
if f > *m {
k = format!(">{}",m);
}
fill_stati(&mut sta,&k,&self.r.cells[i]);
}
let mut g1 = 0.0_f32;
for child in sta.child[0].child.iter(){
g1 += (child.n as f32 / sta.child[0].n as f32) * (child.n as f32 / sta.child[0].n as f32);
}
g1 = 1.0 - g1;
let mut g2 = 0.0_f32;
for child in sta.child[1].child.iter(){
g2 += (child.n as f32 / sta.child[1].n as f32) * (child.n as f32 / sta.child[1].n as f32);
}
g2 = 1.0 - g2;
let g = (sta.child[0].n as f32/sta.n as f32) * g1 + (sta.child[1].n as f32/sta.n as f32) * g2;
if g < min_gini{
min_gini = g;
devide_val = *m;
}
}
c.min_gini = min_gini;
c.devide_val = devide_val;
}
}
}
fn get_discr_stati(&mut self){
for c in self.cols.iter_mut(){
if let ColType::Discr = c.t {
let mut sta = Stati::new(c.name.clone());
for i in 0..c.cells.len() {
let k = &c.cells[i];
sta.n += 1;
fill_stati(&mut sta,k,&self.r.cells[i]);
}
c.stati = sta;
}
}
}
fn get_discr_min_gini_and_sel_val(&mut self){
for c in self.cols.iter_mut(){
if let ColType::Discr = c.t {
let mut min_gini = f32::MAX;
let mut sel_val = String::from("");
for sta in c.stati.child.iter() {
let mut g1 = 0.0_f32;
for child in sta.child.iter(){
g1 += (child.n as f32 / sta.n as f32) * (child.n as f32 / sta.n as f32);
}
g1 = 1.0 - g1;
let mut sta_buff = Stati::new("others".to_string());
for sta2 in c.stati.child.iter(){
if sta2.k == sta.k{
continue;
}
for sta2_child in sta2.child.iter(){
sta_buff.n += sta2_child.n;
let mut has = false;
for buff_child in sta_buff.child.iter_mut(){
if buff_child.k == sta2_child.k{
buff_child.n += sta2_child.n;
has = true;
break;
}
}
if !has {
let mut n_sta = Stati::new(sta2_child.k.clone());
n_sta.n = sta2_child.n;
sta_buff.child.push(n_sta);
}
}
}
let mut g2 = 0.0_f32;
for child in sta_buff.child.iter(){
g2 += (child.n as f32 / sta_buff.n as f32) * (child.n as f32 / sta_buff.n as f32);
}
g2 = 1.0 - g2;
let g = (sta.n as f32/c.stati.n as f32) * g1 + (sta_buff.n as f32/c.stati.n as f32) * g2;
if g < min_gini{
min_gini = g;
sel_val = sta.k.clone();
}
}
c.min_gini = min_gini;
c.sel_val = sel_val;
}
}
}
pub fn push(&mut self,v:Vec<String>){
for s in v.iter(){
if s == ""{
return;
}
}
for i in 0..(v.len()-1) {
self.cols[i].cells.push(v[i].clone());
}
self.r.cells.push(v[v.len()-1].clone());
}
pub fn gen_col(&mut self,name:Vec<String>,t:Vec<ColType>){
for i in 0..name.len() {
let c = Col::new(name[i].clone(),t[i]);
self.cols.push(c);
}
}
pub fn new()->D{
D{
cols:Vec::new(),
r:Result::new(),
}
}
}
#[derive(Debug)]
pub enum NodeType {
Name,
Value,
}
#[derive(Debug)]
pub struct Node {
name:String,
child:Vec<Node>,
t:NodeType,
}
impl Node{
fn new(name:String,t:NodeType)->Node{
Node{
name,
child:Vec::new(),
t,
}
}
}
pub fn test_by_tree(node:&Node,name:&Vec<String>,t:&Vec<ColType>,v:&Vec<String>)->String{
let mut cur_node = node;
loop{
if cur_node.child.len() == 0{
break;
}
match cur_node.t{
NodeType::Name=>{
for i in 0..name.len(){
if name[i] == cur_node.name{
match t[i]{
ColType::Conti=>{
let vi:f32 = v[i].parse().unwrap();
let mut ncv = cur_node.child[0].name.clone();
ncv.remove(0);
let vn:f32 = ncv.parse().unwrap();
for c_n in cur_node.child.iter(){
if c_n.name.starts_with("<"){
if vi < vn{
cur_node = c_n;
break;
}
}else{
if vi > vn{
cur_node = c_n;
break;
}
}
}
},
ColType::Discr=>{
for c_n in cur_node.child.iter(){
if c_n.name == v[i]{
cur_node = c_n;
break;
}else if c_n.name == "!!!"{
cur_node = c_n;
break;
}
}
},
}
break;
}
}
},
NodeType::Value=>{
cur_node = cur_node.child.first().unwrap();
},
}
}
cur_node.name.clone()
}
pub fn get_tree(mut d:D,times:u8,max_times:u8)->Option<Node>{
if d.r.cells.len() == 0{
return None;
}
let result = d.r.cells[0].clone();
let mut has_diff = false;
for i in 1..d.r.cells.len(){
if result != d.r.cells[i]{
has_diff = true;
break;
}
}
if !has_diff{
let nd = Node::new(d.r.cells[0].clone(),NodeType::Value);
return Some(nd);
}
if times > max_times{
let mut sta = Stati::new("sdjfs".to_string());
for s in d.r.cells.iter(){
fill_stati(&mut sta,s,&String::from(""));
}
let mut max = 0;
for c in sta.child.iter(){
if c.n > max{
max = c.n;
}
}
for c in sta.child.iter(){
if c.n == max{
let nd = Node::new(c.k.clone(),NodeType::Value);
return Some(nd);
}
}
}
d.get_discr_stati();
d.get_discr_min_gini_and_sel_val();
d.get_conti_min_gini_and_devide_val();
let mut gr = f32::MAX;
for i in d.cols.iter(){
if i.min_gini < gr{
gr = i.min_gini;
}
}
let mut node = None ;
let mut m_i = 0;
for i in 0..d.cols.len(){
if d.cols[i].min_gini == gr{
m_i = i;
let mut nd = Node::new(d.cols[i].name.clone(),NodeType::Name);
match d.cols[i].t{
ColType::Conti=>{
let k1 = format!("<{}",d.cols[i].devide_val);
let n1 = Node::new(k1,NodeType::Value);
nd.child.push(n1);
let k2 = format!(">{}",d.cols[i].devide_val);
let n2 = Node::new(k2,NodeType::Value);
nd.child.push(n2);
},
ColType::Discr=>{
let k1 = d.cols[i].sel_val.clone();
let n1 = Node::new(k1,NodeType::Value);
nd.child.push(n1);
let k2 = "!!!".to_string();
let n2 = Node::new(k2,NodeType::Value);
nd.child.push(n2);
},
}
node = Some(nd);
break;
}
}
if let Some(mut nd) = node{
let mut v_d = Vec::new(); for _j in 0..nd.child.len(){
let di = D::new();
v_d.push(di);
}
if nd.child.len() > 0{
loop{
if d.cols[0].cells.len() == 0 {
break;
}
let curr_col_val = d.cols[m_i].cells.remove(0); for ni in 0..nd.child.len(){
let mut right = false;
match d.cols[m_i].t{
ColType::Conti=>{
let mut cp = nd.child[ni].name.clone();
cp.remove(0);
let f1:f32 = cp.parse().unwrap();
let f2:f32 = curr_col_val.parse().unwrap();
if nd.child[ni].name.starts_with(">"){
if f2 > f1{
right = true;
}
}else{
if f2 < f1{
right = true;
}
}
},
ColType::Discr=>{
if nd.child[ni].name == curr_col_val{
right = true;
}else if nd.child[ni].name == "!!!"{
right = true;
}
},
}
if right{
for i in 0..d.cols.len(){
let val;
if i == m_i{
val = curr_col_val.clone();
}else{
val = d.cols[i].cells.remove(0);
}
let mut has = false;
for ac in v_d[ni].cols.iter_mut(){
if ac.name == d.cols[i].name {
ac.cells.push(val.clone());
has = true;
break;
}
}
if !has{
let mut col = Col::new(d.cols[i].name.clone(),d.cols[i].t);
col.cells.push(val);
v_d[ni].cols.push(col);
}
}
let r = d.r.cells.remove(0);
v_d[ni].r.cells.push(r);
break;
}
}
}
}
for i in nd.child.iter_mut(){
let di = v_d.remove(0);
if let Some(n) = get_tree(di,times + 1,max_times){
i.child.push(n);
}
}
return Some(nd);
}
None
}
impl fmt::Display for Node{
fn fmt(&self,f:&mut fmt::Formatter)->fmt::Result{
write_node(self,f,0)
}
}
fn write_node(n:&Node,f:&mut fmt::Formatter,layer:usize)->fmt::Result{
for _i in 0..layer{
write!(f,"\t|")?;
}
write!(f,"-{}\r\n",if n.name == "!!!"{"其他".to_string()}else{n.name.clone()})?;
if n.child.len()>0{
for nd in n.child.iter(){
write_node(nd,f,layer+1)?
}
}
write!(f,"")
}
fn fill_stati(sta:&mut Stati,k1:&String,k2:&String) {
let mut has = false;
for mut ci in sta.child.iter_mut() {
if ci.k == *k1 {
ci.n += 1;
has = true;
if k2 != "" {
fill_stati(&mut ci,k2,&"".to_string());
}
break;
}
}
if !has {
let mut s = Stati::new(k1.clone());
s.n += 1;
if k2 != "" {
let mut sc = Stati::new(k2.clone());
sc.n += 1;
s.child.push(sc);
}
sta.child.push(s);
}
}
fn sort(v:&mut Vec<f32>) -> Vec<f32> {
let mut r = Vec::new();
loop {
if v.len() == 0{
break;
}
let mut min = v[0];
for i in 1..v.len(){
min = min.min(v[i]);
}
for i in 0..v.len(){
if min == v[i]{
r.push(v.remove(i));
break;
}
}
}
r
}