Skip to main content

qubic_engine/
ai.rs

1#![feature(alloc, heap_api)]
2#![feature(ptr_internals)]
3#[allow(warnings)]
4pub mod line;
5#[cfg(not(target_arch = "wasm32"))]
6pub mod line_nn;
7pub mod mcts;
8pub mod mpc;
9#[cfg(not(target_arch = "wasm32"))]
10pub mod pattern;
11#[cfg(not(target_arch = "wasm32"))]
12pub mod position;
13pub mod timeout;
14
15use super::board::{Board, GetAction};
16use super::ml::{Graph, Tensor};
17use crate::board::{
18    self, count_1row, count_2row, count_3row, get_random, get_reach_mask, mate_check_horizontal,
19    pprint_board,
20};
21use crate::utills::rand::get_random_usize;
22// use ort::{Environment, GraphOptimizationLevel, Session, SessionBuilder};
23use anyhow::{Ok, Result};
24use serde::ser::SerializeStruct;
25use serde::{Deserialize, Serialize};
26use std::cell::RefCell;
27use std::collections::HashMap;
28use std::ptr::NonNull;
29use std::sync::mpsc::RecvTimeoutError;
30use std::sync::{Arc, Mutex};
31use std::time::Instant;
32use std::{f32, thread};
33
34pub const MAX: i32 = 1600;
35const F32_INVERSE_LAMBDA: f32 = 0.9999;
36const F32_INVERSE_LAMBDA_INVERSE: f32 = 20000.0 / 19999.0;
37const F32_INVERSE_BIAS: f32 = 0.99995;
38
39pub fn negmax<F>(b: &Board, depth: u8, eval_func: &F) -> (u8, i32, i32)
40where
41    F: Fn(&Board) -> i32,
42{
43    let mut count = 0;
44    let mut max_val = -MAX - 1;
45    let max_action = 16;
46    let mut max_action: u8 = 16;
47    let actions = b.valid_actions();
48    for action in actions.iter() {
49        let next_board = &b.next(*action);
50        if next_board.is_win() {
51            return (*action, -MAX, count);
52        } else if next_board.is_draw() {
53            return (*action, 0, count);
54        } else if depth <= 1 {
55            let val = eval_func(b);
56            if max_val < val {
57                max_val = val;
58                max_action = *action;
59            }
60        } else {
61            let (_, val, _count) = negmax(next_board, depth - 1, eval_func);
62            count += 1 + _count;
63            if max_val < val {
64                max_val = val;
65                max_action = *action;
66            }
67        }
68    }
69
70    return (max_action, max_val, count);
71}
72
73pub fn negalpha(
74    b: &Board,
75    depth: u8,
76    alpha: i32,
77    beta: i32,
78    e: &Box<dyn Evaluator>,
79) -> (u8, i32, i32) {
80    // println!("depth:{depth}, alpha:{alpha}, beta:{beta}");
81    // pprint_board(b);
82    let mut count = 0;
83    let actions = b.valid_actions();
84    let mut max_val = -MAX - 1;
85    let mut max_actions = Vec::new();
86    let mut alpha = alpha;
87
88    if depth <= 1 {
89        for action in actions.iter() {
90            let next_board = &b.next(*action);
91            if next_board.is_win() {
92                return (*action, MAX, count);
93            } else if next_board.is_draw() {
94                return (*action, 0, count);
95            }
96            let val = -e.eval_func(next_board);
97            if max_val < val {
98                max_val = val;
99                max_actions = vec![*action];
100                if max_val > alpha {
101                    alpha = max_val;
102                    if alpha > beta {
103                        // println!("[{depth}]->max_val:{max_val}");
104                        return (*action, max_val, count);
105                    }
106                }
107            } else if max_val == val {
108                max_actions.push(*action);
109            }
110        }
111    } else {
112        let mut action_nb_vals: Vec<(u8, Board, i32)> = Vec::new();
113
114        for action in actions.into_iter() {
115            let next_board = b.next(action);
116            if next_board.is_win() {
117                return (action, MAX, count);
118            } else if next_board.is_draw() {
119                return (action, 0, count);
120            }
121
122            let val = -e.eval_func(&next_board);
123            action_nb_vals.push((action, next_board, val));
124        }
125
126        action_nb_vals.sort_by(|a, b| a.2.cmp(&b.2).reverse());
127        // for (a, b, c) in action_nb_vals.iter() {
128        //     print!("[{}, {}]", a, c);
129        // }
130        // println!("");
131
132        for (action, next_board, val) in action_nb_vals {
133            let (_, val, _count) = negalpha(&next_board, depth - 1, -beta, -alpha, e);
134            count += 1 + _count;
135            let val = -999 * val / 1000;
136            if max_val < val {
137                max_val = val;
138                max_actions = vec![action];
139                if max_val > alpha {
140                    alpha = max_val;
141                    if alpha > beta {
142                        // println!("[{depth}]->max_val:{max_val}");
143                        return (action, max_val, count);
144                    }
145                }
146            } else if max_val == val {
147                max_actions.push(action);
148            }
149        }
150    }
151    return (
152        max_actions[get_random_usize() % max_actions.len()],
153        max_val,
154        count,
155    );
156}
157
158pub fn negalphaf(
159    b: &Board,
160    depth: u8,
161    alpha: f32,
162    beta: f32,
163    e: &Box<dyn EvaluatorF>,
164) -> (u8, f32, i32) {
165    //
166    let mut count = 0;
167    let actions = b.valid_actions();
168    let mut max_val = -2.0;
169    let mut max_actions = Vec::new();
170    let mut alpha = alpha;
171
172    if depth <= 1 {
173        for action in actions.iter() {
174            let next_board = &b.next(*action);
175            if next_board.is_win() {
176                return (*action, 1.0, count);
177            } else if next_board.is_draw() {
178                return (*action, 0.5, count);
179            }
180            let val = 1.0 - e.eval_func_f32(next_board);
181            if max_val < val {
182                max_val = val;
183                max_actions = vec![*action];
184                if max_val > alpha {
185                    alpha = max_val;
186                    if alpha > beta {
187                        // println!("[{depth}]->max_val:{max_val}");
188                        return (*action, max_val, count);
189                    }
190                }
191            } else if max_val == val {
192                max_actions.push(*action);
193            }
194        }
195    } else {
196        let mut action_nb_vals: Vec<(u8, Board, f32)> = Vec::new();
197
198        for action in actions.into_iter() {
199            let next_board = b.next(action);
200            if next_board.is_win() {
201                return (action, 1.0, count);
202            } else if next_board.is_draw() {
203                return (action, 0.5, count);
204            }
205
206            let val = 1.0 - e.eval_func_f32(&next_board);
207            action_nb_vals.push((action, next_board, val));
208        }
209
210        action_nb_vals.sort_by(|a, b| {
211            // a.2.cmp(&b.2).reverse();
212            b.2.partial_cmp(&a.2).unwrap()
213        });
214
215        for (action, next_board, _) in action_nb_vals {
216            let (_, val, _count) = negalphaf(&next_board, depth - 1, 1.0 - beta, 1.0 - alpha, e);
217            count += 1 + _count;
218            let val = 0.9995 - 0.999 * val;
219            if max_val < val {
220                max_val = val;
221                max_actions = vec![action];
222                if max_val > alpha {
223                    alpha = max_val;
224                    if alpha > beta {
225                        // println!("[{depth}]->max_val:{max_val}");
226                        return (action, max_val, count);
227                    }
228                }
229            } else if max_val == val {
230                max_actions.push(action);
231            }
232        }
233    }
234    return (
235        max_actions[get_random_usize() % max_actions.len()],
236        max_val,
237        count,
238    );
239}
240
241#[derive(Clone, Copy)]
242pub enum Fail {
243    High(f32),
244    Low(f32),
245    Ex(f32),
246}
247
248impl Fail {
249    fn inverse(&self) -> Self {
250        use Fail::*;
251        match self {
252            High(beta) => Low(F32_INVERSE_BIAS - F32_INVERSE_LAMBDA * beta),
253            Low(alpha) => High(F32_INVERSE_BIAS - F32_INVERSE_LAMBDA * alpha),
254            Ex(val) => Ex(F32_INVERSE_BIAS - F32_INVERSE_LAMBDA * val),
255        }
256    }
257    fn ininverse(&self) -> Self {
258        use Fail::*;
259        match self {
260            High(beta) => Low((F32_INVERSE_BIAS - beta) * F32_INVERSE_LAMBDA_INVERSE),
261            Low(alpha) => High((F32_INVERSE_BIAS - alpha) * F32_INVERSE_LAMBDA_INVERSE),
262            Ex(val) => Ex((F32_INVERSE_BIAS - val) * F32_INVERSE_LAMBDA_INVERSE),
263        }
264    }
265
266    fn is_fail(&self) -> bool {
267        use Fail::*;
268        match self {
269            Ex(_) => false,
270            _ => true,
271        }
272    }
273
274    fn is_equal(&self, x: f32) -> bool {
275        match self {
276            Fail::Ex(val) => *val == x,
277            _ => false,
278        }
279    }
280
281    fn f32_minus(&self, x: f32) -> Self {
282        use Fail::*;
283
284        match self {
285            Ex(val) => Ex(x - val),
286            High(val) => Low(x - val),
287            Low(val) => High(x - val),
288        }
289    }
290
291    fn get_val(&self) -> f32 {
292        use Fail::*;
293
294        match self {
295            Ex(val) => *val,
296            High(val) => *val,
297            Low(val) => *val,
298        }
299    }
300
301    fn is_fail_low(&self) -> bool {
302        match self {
303            Fail::Low(_) => true,
304            _ => false,
305        }
306    }
307
308    fn is_fail_high(&self) -> bool {
309        match self {
310            Fail::High(_) => true,
311            _ => false,
312        }
313    }
314
315    fn get_exval(&self) -> Option<f32> {
316        match self {
317            Fail::Ex(x) => Some(*x),
318            _ => None,
319        }
320    }
321
322    fn to_string(&self) -> String {
323        use Fail::*;
324        match self {
325            High(x) => format!("High({x:.10})"),
326            Low(x) => format!("Low({x:.10})"),
327            Ex(x) => format!(" Ex({x:.10})"),
328        }
329    }
330}
331
332impl std::fmt::Debug for Fail {
333    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
334        use Fail::*;
335        match self {
336            High(x) => {
337                let _ = write!(f, "High({x:.10})");
338            }
339            Low(x) => {
340                let _ = write!(f, "Low({x:.10})");
341            }
342            Ex(x) => {
343                let _ = write!(f, " Ex({x:.10})");
344            }
345        }
346
347        Result::Ok(())
348    }
349}
350
351fn print_blank(n: u8) {
352    for i in 0..n {
353        print!(" ");
354    }
355}
356
357pub fn negalphaf_hash(
358    b: &Board,
359    depth: u8,
360    alpha: f32,
361    beta: f32,
362    hashmap: &mut HashMap<u128, Fail>,
363    e: &Box<dyn EvaluatorF>,
364) -> (u8, Fail, i32) {
365    use Fail::*;
366
367    let mut count = 0;
368    let actions = b.valid_actions();
369    let mut max_val = -2.0;
370    let mut max_actions = Vec::new();
371    let mut alpha = alpha;
372
373    // pprint_board(b);
374    // print_blank(5 - depth);
375    // println!("[depth:{depth}]alpha:{alpha}, beta:{beta}");
376
377    if depth <= 1 {
378        for action in actions.iter() {
379            let next_board = &b.next(*action);
380            // let hash = next_board.hash();
381            let hash = b2u128(next_board);
382            let map_val = hashmap.get(&hash);
383            let val;
384            match map_val {
385                Some(&old_val) => {
386                    if old_val.is_equal(0.0) {
387                        return (*action, Ex(1.0), count);
388                    }
389                    val = 1.0 - old_val.get_val();
390                }
391                None => {
392                    if next_board.is_win() {
393                        hashmap.insert(hash, Ex(0.0));
394                        return (*action, Ex(1.0), count);
395                    } else if next_board.is_draw() {
396                        hashmap.insert(hash, Ex(0.5));
397                        return (*action, Ex(0.5), count);
398                    }
399                    let next_val = e.eval_func_f32(next_board);
400                    val = 1.0 - next_val;
401                    hashmap.insert(hash, Ex(next_val));
402                }
403            }
404            if max_val < val {
405                max_val = val;
406                max_actions = vec![*action];
407                if max_val > alpha {
408                    alpha = max_val;
409                    if alpha > beta {
410                        // println!("[{depth}]->max_val:{max_val}");
411                        return (*action, High(max_val), count);
412                    }
413                }
414            } else if max_val == val {
415                max_actions.push(*action);
416            }
417        }
418    } else {
419        let mut action_nb_vals: Vec<(u8, Board, f32, u128, Option<Fail>)> = Vec::new();
420
421        for action in actions.into_iter() {
422            let next_board = b.next(action);
423            // let hash = next_board.hash();
424            let hash = b2u128(&next_board);
425            let map_val = hashmap.get(&hash);
426
427            match map_val {
428                Some(&old_val) => {
429                    if old_val.is_equal(0.0) {
430                        return (action, Ex(1.0), count);
431                    }
432                    action_nb_vals.push((
433                        action,
434                        next_board,
435                        old_val.f32_minus(1.0).get_val(),
436                        hash,
437                        Some(old_val.inverse()),
438                    ));
439                }
440                None => {
441                    if next_board.is_win() {
442                        hashmap.insert(hash, Ex(0.0));
443                        // print_blank(5 - depth);
444                        // println!("[depth:{depth}]action:{action}, win");
445                        return (action, Ex(1.0), count);
446                    } else if next_board.is_draw() {
447                        hashmap.insert(hash, Ex(0.5));
448                        // print_blank(5 - depth);
449                        // println!("[depth:{depth}]action:{action}, draw");
450                        return (action, Ex(0.5), count);
451                    }
452                    let val = 1.0 - e.eval_func_f32(&next_board);
453                    action_nb_vals.push((action, next_board, val, hash, None));
454                }
455            }
456        }
457
458        action_nb_vals.sort_by(|a, b| {
459            // a.2.cmp(&b.2).reverse();
460            b.2.partial_cmp(&a.2).unwrap()
461        });
462
463        for (action, next_board, old_val, hash, hit) in action_nb_vals {
464            let val;
465            // println!("[depth:{depth}], action:{action}, alpha:{alpha}, beta:{beta}",);
466            if let Some(fail_val) = hit {
467                // if depth > 2 {
468                //     println!("[{depth}]hit, {}", fail_val.to_string());
469                // }
470                match fail_val {
471                    High(x) => {
472                        if beta < x {
473                            return (action, High(x), count);
474                        } else {
475                            let new_alpha = x.max(alpha);
476                            let (_, _val, _count) = negalphaf_hash(
477                                &next_board,
478                                depth - 1,
479                                1.0 - beta,
480                                1.0 - new_alpha,
481                                hashmap,
482                                e,
483                            );
484                            count += _count;
485                            hashmap.insert(hash, _val);
486                            let _val = _val.inverse();
487                            if _val.is_fail_high() {
488                                return (action, High(beta), count);
489                            }
490                            val = _val.get_val();
491                        }
492                    }
493                    Low(x) => {
494                        if alpha > x {
495                            continue;
496                        } else {
497                            let new_beta = x.min(beta);
498                            // fial_low(alpha) or fail_ex(val) < new_beta
499                            let (_, _val, _count) = negalphaf_hash(
500                                &next_board,
501                                depth - 1,
502                                1.0 - new_beta,
503                                1.0 - alpha,
504                                hashmap,
505                                e,
506                            );
507                            hashmap.insert(hash, _val);
508                            let _val = _val.inverse();
509                            if _val.is_fail_low() {
510                                continue;
511                            }
512                            val = _val.get_val();
513                        }
514                    }
515                    Ex(x) => {
516                        // val = F32_INVERSE_BIAS - F32_INVERSE_BIAS * x;
517                        val = x;
518                    }
519                }
520            } else {
521                let (_, _val, _count) =
522                    negalphaf_hash(&next_board, depth - 1, 1.0 - beta, 1.0 - alpha, hashmap, e);
523                hashmap.insert(hash, _val);
524                count += 1 + _count;
525                let _val = _val.inverse();
526
527                match _val {
528                    High(x) => return (action, High(x), count),
529                    Low(_) => continue,
530                    Ex(x) => {
531                        val = x;
532                    }
533                }
534            }
535            if max_val < val {
536                max_val = val;
537                max_actions = vec![action];
538                if max_val > alpha {
539                    alpha = max_val;
540                    if alpha > beta {
541                        // println!("[{depth}]->max_val:{max_val}");
542                        return (action, High(max_val), count);
543                    }
544                }
545            } else if max_val == val {
546                max_actions.push(action);
547            }
548        }
549    }
550    if max_actions.len() == 0 {
551        return (201, Low(alpha), count);
552    }
553
554    return (
555        max_actions[get_random_usize() % max_actions.len()],
556        Ex(max_val),
557        count,
558    );
559}
560
561pub fn negalphaf_hash_iter(
562    b: &Board,
563    depth: u8,
564    alpha: f32,
565    beta: f32,
566    gen: u8,
567    hashmap: &mut HashMap<u128, (Fail, u8)>,
568    e: &Box<dyn EvaluatorF>,
569    top: bool,
570) -> (u8, Fail, i32) {
571    use Fail::*;
572
573    let mut count = 0;
574    let actions = b.valid_actions();
575    let mut max_val = -2.0;
576    let mut max_actions = Vec::new();
577    let mut alpha = alpha;
578
579    // pprint_board(b);
580    // print_blank(5 - depth);
581    // println!("[depth:{depth}]alpha:{alpha}, beta:{beta}");
582
583    if depth <= 1 {
584        for action in actions.iter() {
585            let next_board = &b.next(*action);
586            // let hash = next_board.hash();
587            let hash = b2u128(next_board);
588            let map_val = hashmap.get(&hash);
589            let val;
590            match map_val {
591                Some(&(old_val, old_gen)) => {
592                    if old_val.is_equal(0.0) {
593                        return (*action, Ex(1.0), count);
594                    } else if next_board.is_draw() {
595                        return (*action, Ex(0.5), count);
596                    }
597                    val = 1.0 - old_val.get_val();
598                }
599                None => {
600                    if next_board.is_win() {
601                        hashmap.insert(hash, (Ex(0.0), gen));
602                        return (*action, Ex(1.0), count);
603                    } else if next_board.is_draw() {
604                        hashmap.insert(hash, (Ex(0.5), gen));
605                        return (*action, Ex(0.5), count);
606                    }
607                    let next_val = e.eval_func_f32(next_board);
608                    val = 1.0 - next_val;
609                    hashmap.insert(hash, (Ex(next_val), gen));
610                }
611            }
612            if max_val < val {
613                max_val = val;
614                max_actions = vec![*action];
615                if max_val > alpha {
616                    alpha = max_val;
617                    if alpha > beta {
618                        // println!("[{depth}]->max_val:{max_val}");
619                        return (*action, High(max_val), count);
620                    }
621                }
622            } else if max_val == val {
623                max_actions.push(*action);
624            }
625        }
626    } else {
627        let mut action_nb_vals: Vec<(u8, Board, f32, u128, (Option<Fail>, u8))> = Vec::new();
628
629        for action in actions.into_iter() {
630            let next_board = b.next(action);
631            // let hash = next_board.hash();
632            let hash = b2u128(&next_board);
633            let map_val = hashmap.get(&hash);
634
635            match map_val {
636                Some(&(old_val, old_gen)) => {
637                    if old_val.is_equal(0.0) {
638                        return (action, Ex(1.0), count);
639                    } else if next_board.is_draw() {
640                        return (action, Ex(0.5), count);
641                    }
642                    action_nb_vals.push((
643                        action,
644                        next_board,
645                        old_val.f32_minus(1.0).get_val(),
646                        hash,
647                        (Some(old_val.inverse()), old_gen),
648                    ));
649                }
650                None => {
651                    if next_board.is_win() {
652                        hashmap.insert(hash, (Ex(0.0), gen));
653                        // print_blank(5 - depth);
654                        // println!("[depth:{depth}]action:{action}, win");
655                        return (action, Ex(1.0), count);
656                    } else if next_board.is_draw() {
657                        hashmap.insert(hash, (Ex(0.5), gen));
658                        // print_blank(5 - depth);
659                        // println!("[depth:{depth}]action:{action}, draw");
660                        return (action, Ex(0.5), count);
661                    }
662                    let val = 1.0 - e.eval_func_f32(&next_board);
663                    action_nb_vals.push((action, next_board, val, hash, (None, 0)));
664                }
665            }
666        }
667
668        action_nb_vals.sort_by(|a, b| {
669            // a.2.cmp(&b.2).reverse();
670            b.2.partial_cmp(&a.2).unwrap()
671        });
672
673        for (action, next_board, old_val, hash, (hit, old_gen)) in action_nb_vals {
674            let val;
675            if old_gen == gen {
676                if let Some(fail_val) = hit.clone() {
677                    // if depth > 2 {
678                    //     println!("[{depth}]hit, {}", fail_val.to_string());
679                    // }
680                    match fail_val {
681                        High(x) => {
682                            if beta < x {
683                                return (action, High(x), count);
684                            } else {
685                                let new_alpha = x.max(alpha);
686                                let (_, _val, _count) = negalphaf_hash_iter(
687                                    &next_board,
688                                    depth - 1,
689                                    1.0 - beta,
690                                    1.0 - new_alpha,
691                                    gen,
692                                    hashmap,
693                                    e,
694                                    false,
695                                );
696                                count += _count;
697                                hashmap.insert(hash, (_val, gen));
698                                let _val = _val.inverse();
699                                if _val.is_fail_high() {
700                                    return (action, High(beta), count);
701                                }
702                                val = _val.get_val();
703                            }
704                        }
705                        Low(x) => {
706                            if alpha > x {
707                                if cfg!(feature = "view_detail") {
708                                    if top {
709                                        match hit{
710                                            Some(fail)=>println!("action:{action:>2}, val:Low({x}), old:{fail:#?}-{old_val}"),
711                                            None => println!("action:{action:>2}, val:Low({x}), old:None-{old_val}"),
712                                        }
713                                    }
714                                }
715                                continue;
716                            } else {
717                                let new_beta = x.min(beta);
718                                // fial_low(alpha) or fail_ex(val) < new_beta
719                                let (_, _val, _count) = negalphaf_hash_iter(
720                                    &next_board,
721                                    depth - 1,
722                                    1.0 - new_beta,
723                                    1.0 - alpha,
724                                    gen,
725                                    hashmap,
726                                    e,
727                                    false,
728                                );
729                                hashmap.insert(hash, (_val, gen));
730                                let _val = _val.inverse();
731                                if _val.is_fail_low() {
732                                    continue;
733                                }
734                                val = _val.get_val();
735                            }
736                        }
737                        Ex(x) => {
738                            // val = F32_INVERSE_BIAS - F32_INVERSE_BIAS * x;
739                            val = x;
740                        }
741                    }
742                } else {
743                    let (_, _val, _count) = negalphaf_hash_iter(
744                        &next_board,
745                        depth - 1,
746                        1.0 - beta,
747                        1.0 - alpha,
748                        gen,
749                        hashmap,
750                        e,
751                        false,
752                    );
753                    hashmap.insert(hash, (_val, gen));
754                    count += 1 + _count;
755                    let _val = _val.inverse();
756
757                    match _val {
758                        High(x) => return (action, High(x), count),
759                        Low(x) => {
760                            if cfg!(feature = "view_detail") {
761                                if top {
762                                    match hit {
763                                        Some(fail) => println!(
764                                            "action:{action:>2}, val:Low({x:.4}), old:{fail:#?}-{old_val:.4}"
765                                        ),
766                                        None => println!(
767                                            "action:{action:>2}, val:Low({x:.4}), old:None-{old_val:.4}"
768                                        ),
769                                    }
770                                }
771                            }
772                            continue;
773                        }
774                        Ex(x) => {
775                            val = x;
776                        }
777                    }
778                }
779            } else {
780                let (_, _val, _count) = negalphaf_hash_iter(
781                    &next_board,
782                    depth - 1,
783                    1.0 - beta,
784                    1.0 - alpha,
785                    gen,
786                    hashmap,
787                    e,
788                    false,
789                );
790                hashmap.insert(hash, (_val, gen));
791                count += 1 + _count;
792                let _val = _val.inverse();
793
794                match _val {
795                    High(x) => return (action, High(x), count),
796                    Low(x) => {
797                        if cfg!(feature = "view_detail") {
798                            if top {
799                                match hit {
800                                    Some(fail) => println!(
801                                        "action:{action:>2}, val:Low({x:.4}), old:{fail:#?}-{old_val:.4}"
802                                    ),
803                                    None => println!(
804                                        "action:{action:>2}, val:Low({x:.4}), old:None-{old_val:.4}"
805                                    ),
806                                }
807                            }
808                        }
809                        continue;
810                    }
811                    Ex(x) => {
812                        val = x;
813                    }
814                }
815            }
816
817            if cfg!(feature = "view_detail") {
818                if top {
819                    match hit {
820                        Some(x) => {
821                            println!("action:{action:>2}, val:{val:.4}, old:{x:#?}-{old_val:.4}")
822                        }
823                        None => println!("action:{action:>2}, val:{val:.4}, old:None-{old_val:.4}"),
824                    }
825                }
826            }
827
828            if max_val < val {
829                max_val = val;
830                max_actions = vec![action];
831                if max_val > alpha {
832                    alpha = max_val;
833                    if alpha > beta {
834                        return (action, High(max_val), count);
835                    }
836                }
837            } else if max_val == val {
838                max_actions.push(action);
839            }
840        }
841    }
842    if max_actions.len() == 0 {
843        return (201, Low(alpha), count);
844    }
845
846    return (
847        max_actions[get_random_usize() % max_actions.len()],
848        Ex(max_val),
849        count,
850    );
851}
852
853pub fn negscoutf_hash_iter(
854    b: &Board,
855    depth: u8,
856    alpha: f32,
857    beta: f32,
858    gen: u8,
859    hashmap: &mut HashMap<u128, (Fail, u8)>,
860    e: &Box<dyn EvaluatorF>,
861    top: bool,
862) -> (u8, Fail, i32) {
863    use Fail::*;
864
865    let mut count = 0;
866    let actions = b.valid_actions();
867
868    let (att, def) = b.get_att_def();
869    let four_mask = get_reach_mask(att, def);
870    if four_mask != 0 {
871        let action = four_mask.trailing_zeros() % 16;
872        return (action as u8, Ex(1.0), 0);
873    }
874    let stone = att | def;
875    // あと一箇所しか置ける場所がなければ引き分けが確定している。
876    if stone.count_ones() == 63 {
877        let action = !stone.trailing_zeros() % 16;
878        return (action as u8, Ex(0.5), 0);
879    }
880
881    let mut max_val = -2.0;
882    let mut max_actions = Vec::new();
883    let mut alpha = alpha;
884
885    if depth <= 1 {
886        for action in actions.iter() {
887            let next_board = &b.next(*action);
888            let next_val = e.eval_func_f32(next_board);
889            let val = 1.0 - next_val;
890            if max_val < val {
891                max_val = val;
892                max_actions = vec![*action];
893                if max_val > alpha {
894                    alpha = max_val;
895                    if alpha > beta {
896                        return (*action, High(max_val), count);
897                    }
898                }
899            } else if max_val == val {
900                max_actions.push(*action);
901            }
902        }
903    } else {
904        let mut action_nb_vals: Vec<(u8, Board, f32, u128, (Option<Fail>, u8))> = Vec::new();
905
906        for action in actions.into_iter() {
907            let next_board = b.next(action);
908            // let hash = next_board.hash();
909            let hash = b2u128(&next_board);
910            let map_val = hashmap.get(&hash);
911
912            match map_val {
913                Some(&(old_val, old_gen)) => {
914                    action_nb_vals.push((
915                        action,
916                        next_board,
917                        // old_val.f32_minus(1.0).get_val(),
918                        match old_val {
919                            Ex(x) => 1.0 - x,
920                            Low(x) => -x,
921                            High(x) => -x,
922                        },
923                        hash,
924                        (Some(old_val.inverse()), old_gen),
925                    ));
926                }
927                None => {
928                    let val = 1.0 - e.eval_func_f32(&next_board);
929                    action_nb_vals.push((action, next_board, val, hash, (None, 0)));
930                }
931            }
932        }
933
934        action_nb_vals.sort_by(|a, b| {
935            // a.2.cmp(&b.2).reverse();
936            b.2.partial_cmp(&a.2).unwrap()
937        });
938
939        for (idx, &(action, ref next_board, old_val, hash, (hit, old_gen))) in
940            action_nb_vals.iter().enumerate()
941        {
942            let val;
943            if old_gen == gen {
944                if let Some(fail_val) = hit.clone() {
945                    match fail_val {
946                        High(x) => {
947                            if beta < x {
948                                return (action, High(x), count);
949                            } else {
950                                let new_alpha = x.max(alpha);
951                                let (_, _val, _count) = negscoutf_hash_iter(
952                                    &next_board,
953                                    depth - 1,
954                                    1.0 - beta,
955                                    1.0 - new_alpha,
956                                    gen,
957                                    hashmap,
958                                    e,
959                                    false,
960                                );
961                                count += _count;
962                                hashmap.insert(hash, (_val, gen));
963                                let _val = _val.inverse();
964                                if _val.is_fail_high() {
965                                    return (action, High(beta), count);
966                                }
967                                val = _val.get_val();
968                            }
969                        }
970                        Low(x) => {
971                            if alpha > x {
972                                if cfg!(feature = "view_detail") {
973                                    if top {
974                                        match hit{
975                                            Some(fail)=>println!("action:{action:>2}, val:Low({x}), old:{fail:#?}-{old_val}"),
976                                            None => println!("action:{action:>2}, val:Low({x}), old:None-{old_val}"),
977                                        }
978                                    }
979                                }
980                                continue;
981                            } else {
982                                // null window search
983                                // if depth == 7 && idx > 0{
984                                //     let (_, _val, _) = negscoutf_hash_iter(
985                                //         &next_board,
986                                //         depth - 1,
987                                //         (F32_INVERSE_BIAS - alpha) * F32_INVERSE_LAMBDA_INVERSE - 0.00002,
988                                //         (F32_INVERSE_BIAS - alpha) * F32_INVERSE_LAMBDA_INVERSE - 0.00001,
989                                //         gen,
990                                //         hashmap,
991                                //         e,
992                                //         false,
993                                //     );
994                                //     // println!("{_val:#?}");
995
996                                //     if _val.is_fail_high() {
997                                //         continue;
998                                //     }
999                                // }
1000
1001                                let new_beta = x.min(beta);
1002                                // fial_low(alpha) or fail_ex(val) < new_beta
1003                                let (_, _val, _count) = negscoutf_hash_iter(
1004                                    &next_board,
1005                                    depth - 1,
1006                                    1.0 - new_beta,
1007                                    1.0 - alpha,
1008                                    gen,
1009                                    hashmap,
1010                                    e,
1011                                    false,
1012                                );
1013                                hashmap.insert(hash, (_val, gen));
1014                                let _val = _val.inverse();
1015                                if _val.is_fail_low() {
1016                                    continue;
1017                                }
1018                                val = _val.get_val();
1019                            }
1020                        }
1021                        Ex(x) => {
1022                            // val = F32_INVERSE_BIAS - F32_INVERSE_BIAS * x;
1023                            val = x;
1024                        }
1025                    }
1026                } else {
1027                    // if depth == 7 && idx > 0{
1028                    //     let (_, _val, _) = negscoutf_hash_iter(
1029                    //         &next_board,
1030                    //         depth - 1,
1031                    //         (F32_INVERSE_BIAS - alpha) * F32_INVERSE_LAMBDA_INVERSE - 0.00002,
1032                    //         (F32_INVERSE_BIAS - alpha) * F32_INVERSE_LAMBDA_INVERSE - 0.00001,
1033                    //         gen,
1034                    //         hashmap,
1035                    //         e,
1036                    //         false,
1037                    //     );
1038                    //     // println!("{_val:#?}");
1039                    //     if _val.is_fail_high(){
1040                    //         continue;
1041                    //     }
1042                    // }
1043
1044                    let (_, _val, _count) = negscoutf_hash_iter(
1045                        &next_board,
1046                        depth - 1,
1047                        1.0 - beta,
1048                        1.0 - alpha,
1049                        gen,
1050                        hashmap,
1051                        e,
1052                        false,
1053                    );
1054                    hashmap.insert(hash, (_val, gen));
1055                    count += 1 + _count;
1056                    let _val = _val.inverse();
1057
1058                    match _val {
1059                        High(x) => return (action, High(x), count),
1060                        Low(x) => {
1061                            if cfg!(feature = "view_detail") {
1062                                if top {
1063                                    match hit {
1064                                        Some(fail) => println!(
1065                                            "action:{action:>2}, val:Low({x:.10}), old:{fail:#?}-{old_val:.4}"
1066                                        ),
1067                                        None => println!(
1068                                            "action:{action:>2}, val:Low({x:.10}), old:None-{old_val:.4}"
1069                                        ),
1070                                    }
1071                                }
1072                            }
1073                            continue;
1074                        }
1075                        Ex(x) => {
1076                            val = x;
1077                        }
1078                    }
1079                }
1080            } else {
1081                // if depth == 7 && idx > 0{
1082                //     let (_, _val, _) = negscoutf_hash_iter(
1083                //         &next_board,
1084                //         depth - 1,
1085                //         (F32_INVERSE_BIAS - alpha) * F32_INVERSE_LAMBDA_INVERSE - 0.00002,
1086                //         (F32_INVERSE_BIAS - alpha) * F32_INVERSE_LAMBDA_INVERSE - 0.00001,
1087                //         gen,
1088                //         hashmap,
1089                //         e,
1090                //         false,
1091                //     );
1092                //     // println!("{_val:#?}");
1093
1094                //     if _val.is_fail_high() {
1095                //         continue;
1096                //     }
1097                // }
1098
1099                let (_, _val, _count) = negscoutf_hash_iter(
1100                    &next_board,
1101                    depth - 1,
1102                    1.0 - beta,
1103                    1.0 - alpha,
1104                    gen,
1105                    hashmap,
1106                    e,
1107                    false,
1108                );
1109                hashmap.insert(hash, (_val, gen));
1110                count += 1 + _count;
1111                let _val = _val.inverse();
1112
1113                match _val {
1114                    High(x) => return (action, High(x), count),
1115                    Low(x) => {
1116                        if cfg!(feature = "view_detail") {
1117                            if top {
1118                                match hit {
1119                                    Some(fail) => println!(
1120                                        "action:{action:>2}, val:Low({x:.10}), old:{fail:#?}-{old_val:.4}"
1121                                    ),
1122                                    None => println!(
1123                                        "action:{action:>2}, val:Low({x:.10}), old:None-{old_val:.4}"
1124                                    ),
1125                                }
1126                            }
1127                        }
1128                        continue;
1129                    }
1130                    Ex(x) => {
1131                        val = x;
1132                    }
1133                }
1134            }
1135
1136            if cfg!(feature = "view_detail") {
1137                if top {
1138                    match hit {
1139                        Some(x) => {
1140                            println!("action:{action:>2}, val:{val:.10}, old:{x:#?}-{old_val:.4}")
1141                        }
1142                        None => {
1143                            println!("action:{action:>2}, val:{val:.10}, old:None-{old_val:.4}")
1144                        }
1145                    }
1146                }
1147            }
1148
1149            if max_val < val {
1150                max_val = val;
1151                max_actions = vec![action];
1152                if max_val > alpha {
1153                    alpha = max_val;
1154                    if alpha > beta {
1155                        return (action, High(max_val), count);
1156                    }
1157                }
1158            } else if max_val == val {
1159                max_actions.push(action);
1160            }
1161        }
1162    }
1163    if max_actions.len() == 0 {
1164        return (201, Low(alpha), count);
1165    }
1166
1167    return (
1168        max_actions[get_random_usize() % max_actions.len()],
1169        Ex(max_val),
1170        count,
1171    );
1172}
1173
1174pub struct NegAlpha {
1175    evaluator: Box<dyn Evaluator>,
1176    depth: u8,
1177}
1178
1179impl NegAlpha {
1180    pub fn new(e: Box<dyn Evaluator>, depth: u8) -> Self {
1181        return NegAlpha {
1182            evaluator: e,
1183            depth: depth,
1184        };
1185    }
1186    pub fn eval_with_negalpha(&self, b: &Board) -> (u8, i32, i32) {
1187        return negalpha(b, self.depth, -MAX - 1, MAX + 1, &self.evaluator);
1188    }
1189}
1190
1191impl GetAction for NegAlpha {
1192    fn get_action(&self, b: &Board) -> u8 {
1193        let start = Instant::now();
1194        let (action, v, count) = negalpha(b, self.depth, -MAX - 1, MAX + 1, &self.evaluator);
1195        // pprint_board(b);
1196        let t = start.elapsed().as_nanos();
1197
1198        if cfg!(feature = "render") {
1199            println!("action:{action}, val:{v}, count:{count}/{t}");
1200        }
1201
1202        return action;
1203    }
1204}
1205
1206pub trait EvalAndAnalyze: EvaluatorF + Analyzer {}
1207pub trait EvalAndActF {
1208    fn eval_and_act(&self, b: &Board) -> (u8, f32);
1209}
1210
1211pub struct NegAlphaF {
1212    evaluator: Box<dyn EvaluatorF>,
1213    depth: u8,
1214    pub hashmap: bool,
1215    pub scout: bool,
1216    pub timelimit: u128,
1217    pub min_depth: u8,
1218}
1219
1220impl NegAlphaF {
1221    pub fn new(e: Box<dyn EvaluatorF>, depth: u8) -> Self {
1222        return NegAlphaF {
1223            evaluator: e,
1224            depth: depth,
1225            hashmap: false,
1226            scout: false,
1227            timelimit: 1000,
1228            min_depth: 1,
1229        };
1230    }
1231
1232    pub fn eval_with_negscout_(&self, b: &Board) -> (u8, f32, i32) {
1233        let limit = Instant::now();
1234
1235        let mut hashmap = HashMap::new();
1236
1237        let (mut action, mut val, mut count);
1238        action = 0;
1239        val = Fail::Ex(0.0);
1240        count = 0;
1241        for i in (1..=self.depth).step_by(2) {
1242            let start = Instant::now();
1243            (action, val, count) =
1244                negscoutf_hash_iter(b, i, -2.0, 2.0, i, &mut hashmap, &self.evaluator, true);
1245            let t = start.elapsed().as_nanos();
1246            if val.get_exval().is_none() {
1247                println!(
1248                    "[depth:{i}], action:{action}, count:{}, time:{},{}",
1249                    hashmap.len(),
1250                    t / 1000000,
1251                    t % 1000000
1252                );
1253                let (att, def) = b.get_att_def();
1254                println!("att:{att}, def:{def}");
1255                pprint_board(b);
1256            }
1257            assert!(
1258                val.get_exval().is_some(),
1259                "[depth:{i}], action:{action}, count:{}, time:{t}",
1260                hashmap.len()
1261            );
1262            if cfg!(feature = "view") {
1263                let val_ = ((1.0 / val.get_exval().unwrap()) - 1.0).ln() * -400.0;
1264                println!(
1265                    "[depth:{i}], action:{action}, val:{:#?}({}), count:{}, time:{t}, rate:{}count/ms",
1266                    val,
1267                    val_ as i32,
1268                    hashmap.len(),
1269                    (count as u128 * 1000000) / t,
1270                );
1271            }
1272            if self.min_depth > i {
1273                continue;
1274            }
1275            if limit.elapsed().as_millis() > self.timelimit {
1276                break;
1277            }
1278        }
1279        if cfg!(feature = "view") {
1280            println!("total_time:{}ms", limit.elapsed().as_millis());
1281        }
1282
1283        return (action, val.get_exval().unwrap(), count);
1284    }
1285    pub fn eval_with_negalpha_(&self, b: &Board) -> (u8, f32, i32) {
1286        let limit = Instant::now();
1287
1288        let mut hashmap = HashMap::new();
1289
1290        let (mut action, mut val, mut count);
1291        action = 0;
1292        val = Fail::Ex(0.0);
1293        count = 0;
1294        let (att, def) = b.get_att_def();
1295        let depth = (((att | def).count_zeros() / 2) * 2 + 1) as u8;
1296        let depth = depth.min(self.depth);
1297        for i in (1..=depth).step_by(2) {
1298            let start = Instant::now();
1299            (action, val, count) =
1300                negalphaf_hash_iter(b, i, -2.0, 2.0, i, &mut hashmap, &self.evaluator, true);
1301            let t = start.elapsed().as_nanos();
1302            if val.get_exval().is_none() {
1303                println!(
1304                    "[depth:{i}], action:{action}, count:{}, time:{},{}",
1305                    hashmap.len(),
1306                    t / 1000000,
1307                    t % 1000000
1308                );
1309                let (att, def) = b.get_att_def();
1310                println!("att:{att}, def:{def}");
1311                pprint_board(b);
1312            }
1313            assert!(
1314                val.get_exval().is_some(),
1315                "[depth:{i}], action:{action}, count:{}, time:{t}",
1316                hashmap.len()
1317            );
1318            if cfg!(feature = "view") {
1319                let val_ = ((1.0 / val.get_exval().unwrap()) - 1.0).ln() * -400.0;
1320                println!(
1321                    "[depth:{i}], action:{action}, val:{:#?}({}), count:{}, time:{t}",
1322                    val,
1323                    val_ as i32,
1324                    hashmap.len()
1325                );
1326            }
1327            if self.min_depth > i {
1328                continue;
1329            }
1330            if limit.elapsed().as_millis() > self.timelimit {
1331                break;
1332            }
1333        }
1334        if cfg!(feature = "view") {
1335            println!("total_time:{}ms", limit.elapsed().as_millis());
1336        }
1337
1338        return (action, val.get_exval().unwrap(), count);
1339    }
1340
1341    pub fn eval_with_negalpha(&self, b: &Board) -> (u8, f32, i32) {
1342        // for &action in actions.iter() {
1343        //     let next_b = b.next(action);
1344        //     let (_, val, _) = negalphaf(&next_b, self.depth - 1, -2.0, 2.0, &self.evaluator);
1345        //     let val = 1.0 - val;
1346        //     println!("[{}]:{}", action, val);
1347        // }
1348        let (att, def) = b.get_att_def();
1349        let stone = (att | def).count_ones() as usize;
1350        if cfg!(feature = "view") {
1351            println!("att:{att}, def:{def}");
1352            let start = Instant::now();
1353            let (action, val, count) = negalphaf(b, self.min_depth, -2.0, 2.0, &self.evaluator);
1354            let t = start.elapsed().as_nanos();
1355            // let start = Instant::now();
1356            // let (action2, val2, count2) = self.eval_with_negalpha_(b);
1357            // let t2 = start.elapsed().as_nanos();
1358            println!("action:{action}, val:{val}, count:{count}/{t}ns",);
1359        }
1360
1361        if self.hashmap {
1362            return self.eval_with_negalpha_(b);
1363        } else if self.scout {
1364            return self.eval_with_negscout_(b);
1365        }
1366        let mut hashmap = HashMap::new();
1367        let (a, b, c) = negalphaf_hash(b, self.depth, -2.0, 2.0, &mut hashmap, &self.evaluator);
1368        if cfg!(feature = "render") {
1369            println!("hashmap_size: {}", hashmap.len());
1370        }
1371        return (a, b.get_exval().unwrap(), c);
1372    }
1373}
1374
1375impl GetAction for NegAlphaF {
1376    fn get_action(&self, b: &Board) -> u8 {
1377        let (action, _, _) = self.eval_with_negalpha(b);
1378        return action;
1379    }
1380}
1381
1382impl EvalAndActF for NegAlphaF {
1383    fn eval_and_act(&self, b: &Board) -> (u8, f32) {
1384        let (action, val, _) = self.eval_with_negalpha(b);
1385        return (action, val);
1386    }
1387}
1388
1389impl EvaluatorF for NegAlphaF {
1390    fn eval_func_f32(&self, b: &Board) -> f32 {
1391        let (_, val, _) = self.eval_with_negalpha(b);
1392        return val;
1393    }
1394}
1395
1396pub trait Evaluator {
1397    fn eval_func(&self, b: &Board) -> i32;
1398}
1399
1400pub trait EvaluatorF {
1401    fn eval_func_f32(&self, b: &Board) -> f32;
1402}
1403
1404pub trait Analyzer {
1405    fn analyze_eval(&self, b: &Board) -> f32;
1406    fn analyze(&self, b: &Board) {
1407        pprint_board(b);
1408        let actions = b.valid_actions();
1409
1410        for &action in actions.iter() {
1411            let next_b = b.next(action);
1412            let val = 1.0 - self.analyze_eval(&next_b);
1413            println!("[{}]:{}", action, val);
1414        }
1415    }
1416}
1417pub struct PositionEvaluator {
1418    posmap: Vec<i32>,
1419}
1420
1421impl PositionEvaluator {
1422    pub fn new(posmap: &[i32]) -> Self {
1423        return PositionEvaluator {
1424            posmap: posmap.to_vec(),
1425        };
1426    }
1427    pub fn simpl(vertex: i32, edge: i32, surface: i32, core: i32) -> Self {
1428        let (v, e, s, c) = (vertex, edge, surface, core);
1429        let posmap = vec![
1430            v, e, e, v, e, s, s, e, e, s, s, e, v, e, e, v, e, s, s, e, s, c, c, s, s, c, c, s, e,
1431            s, s, e, e, s, s, e, s, c, c, s, s, c, c, s, e, s, s, e, v, e, e, v, e, s, s, e, e, s,
1432            s, e, v, e, e, v,
1433        ];
1434        return PositionEvaluator { posmap: posmap };
1435    }
1436
1437    pub fn simpl_alpha(
1438        vertex1: i32,
1439        vertex2: i32,
1440        vertex3: i32,
1441        vertex4: i32,
1442        up_surface: i32,
1443        bt_surfacce: i32,
1444        edge1: i32,
1445        edge2: i32,
1446        edge3: i32,
1447        edge4: i32,
1448        up_core: i32,
1449        bt_core: i32,
1450    ) -> Self {
1451        let (v0, v1, v2, v3, s0, s1, e0, e1, e2, e3, c0, c1) = (
1452            vertex1,
1453            vertex2,
1454            vertex3,
1455            vertex4,
1456            bt_surfacce,
1457            up_surface,
1458            edge1,
1459            edge2,
1460            edge3,
1461            edge4,
1462            bt_core,
1463            up_core,
1464        );
1465
1466        let posmap = vec![
1467            v0, e0, e0, v0, e0, s0, s0, e0, e0, s0, s0, e0, v0, e0, e0, v0, v1, e1, e1, v1, e1, c0,
1468            c0, e1, e1, c0, c0, e1, v1, e1, e1, v1, v2, e2, e2, v2, e2, c1, c1, e2, e2, c1, c1, e2,
1469            v2, e2, e2, v2, v3, e3, e3, v3, e3, s1, s1, e3, e3, s1, s1, e3, v3, e3, e3, v3,
1470        ];
1471
1472        return PositionEvaluator { posmap: posmap };
1473    }
1474
1475    pub fn best() -> Self {
1476        return PositionEvaluator::simpl(6, 1, 5, 8);
1477    }
1478}
1479
1480impl Evaluator for PositionEvaluator {
1481    fn eval_func(&self, b: &Board) -> i32 {
1482        let (mut att, mut def) = b.get_att_def();
1483        let mut val = 0;
1484        for i in 0..64 {
1485            if 1 & att == 1 {
1486                val += self.posmap[i];
1487            } else if 1 & def == 1 {
1488                val -= self.posmap[i]
1489            }
1490            att >>= 1;
1491            def >>= 1;
1492        }
1493        return val;
1494    }
1495}
1496
1497pub struct MLEvaluator {
1498    pub g: Graph,
1499    loss: usize,
1500    g_out: usize,
1501    input: usize,
1502    t: usize,
1503}
1504
1505impl MLEvaluator {
1506    pub fn new(g: Graph) -> Self {
1507        return MLEvaluator {
1508            g: g,
1509            loss: 0,
1510            g_out: 0,
1511            input: 0,
1512            t: 0,
1513        };
1514    }
1515
1516    pub fn default() -> Self {
1517        use super::ml::*;
1518        use super::ml::{funcs::*, optim::*, params::*};
1519        let mut g = Graph::new();
1520        g.optimizer = Some(Box::new(MomentumSGD::new(0.01, 0.9)));
1521        let i1: usize = g.push_placeholder();
1522        let i2: usize = g.push_placeholder();
1523        // let t = g.push_placeholder();
1524
1525        let l1 = Linear::auto(128, 64);
1526        let l1 = g.add_layer(vec![i1], Box::new(l1));
1527
1528        let activate = ClippedReLU::default();
1529        let relu = g.add_layer(vec![l1], Box::new(activate));
1530
1531        let l2 = Linear::auto(64, 16);
1532        let l2 = g.add_layer(vec![relu], Box::new(l2));
1533
1534        let relu2 = g.add_layer(vec![l2], Box::new(ClippedReLU::default()));
1535
1536        let l3 = Linear::auto(16, 1);
1537        let l3 = g.add_layer(vec![relu2], Box::new(l3));
1538
1539        let last = g.add_layer(vec![l3], Box::new(Tanh::new()));
1540
1541        // t = lambda * result + (1 - lambda) * t_in
1542        let loss = g.add_layer(vec![last, i2], Box::new(MSE::new()));
1543
1544        g.set_target(last);
1545        g.set_placeholder(vec![i1]);
1546
1547        return MLEvaluator {
1548            g: g,
1549            loss: loss,
1550            g_out: last,
1551            input: i1,
1552            t: i2,
1553        };
1554    }
1555
1556    pub fn inference(&self, b: &Board) -> f32 {
1557        let (att, def) = b.get_att_def();
1558
1559        let mut att_vec = Vec::new();
1560        let mut def_vec = Vec::new();
1561        for i in 0..64 {
1562            if (att >> i) & 1 == 1 {
1563                att_vec.push(1.0);
1564            } else {
1565                att_vec.push(0.0);
1566            }
1567
1568            if (def >> i) & 1 == 1 {
1569                def_vec.push(1.0);
1570            } else {
1571                def_vec.push(0.0);
1572            }
1573        }
1574
1575        let onehot = [att_vec, def_vec].concat();
1576        let onehot = Tensor::new(onehot, vec![128, 1]);
1577
1578        let val = self.g.inference(vec![onehot]);
1579
1580        return val.get_item().unwrap();
1581    }
1582
1583    pub fn eval_with_negalpha(&self, b: &Board, depth: u8) -> (u8, f32, i32) {
1584        return self.eval_with_negalpha_(b, depth, -2.0, 2.0);
1585    }
1586
1587    pub fn eval_with_negalpha_(
1588        &self,
1589        b: &Board,
1590        depth: u8,
1591        alpha: f32,
1592        beta: f32,
1593    ) -> (u8, f32, i32) {
1594        let mut count = 0;
1595        let actions = b.valid_actions();
1596        let mut max_val = -2.0;
1597        let mut max_action: u8 = 16;
1598        for action in actions.iter() {
1599            let next_board = &b.next(*action);
1600            if next_board.is_win() {
1601                return (*action, 1.0, count);
1602            } else if next_board.is_draw() {
1603                return (*action, 0.0, count);
1604            } else if depth <= 1 {
1605                let val = -self.inference(next_board);
1606                count += 1;
1607                if max_val < val {
1608                    max_val = val;
1609                    max_action = *action;
1610                    if max_val > beta {
1611                        return (max_action, max_val, count);
1612                    }
1613                }
1614            } else {
1615                let (_, val, _count) =
1616                    self.eval_with_negalpha_(next_board, depth - 1, -max_val, -alpha);
1617                let val = -val;
1618                count += _count;
1619                if max_val < val {
1620                    max_val = val;
1621                    max_action = *action;
1622                    if max_val > beta {
1623                        return (max_action, max_val, count);
1624                    }
1625                }
1626            }
1627        }
1628        // println!("[{}]max_val:{}", max_action, max_val);
1629        return (max_action, max_val, count);
1630    }
1631
1632    pub fn train(&mut self) {
1633        self.g.set_placeholder(vec![self.input, self.t]);
1634        self.g.set_target(self.loss);
1635    }
1636
1637    pub fn eval(&mut self) {
1638        self.g.set_placeholder(vec![self.input]);
1639        self.g.set_target(self.g_out);
1640    }
1641
1642    pub fn save(&self, s: String) {
1643        self.g.save(s);
1644    }
1645
1646    pub fn load(&mut self, s: String) {
1647        self.g.load(s);
1648    }
1649}
1650
1651pub fn u2vec(board: u128) -> Vec<f32> {
1652    let mut att_vec = Vec::new();
1653    for i in 0..128 {
1654        if (board >> i) & 1 == 1 {
1655            att_vec.push(1.0);
1656        } else {
1657            att_vec.push(0.0);
1658        }
1659    }
1660    return att_vec;
1661}
1662
1663pub fn onehot_vec(n: usize, idx: usize) -> Vec<f32> {
1664    let mut v = Vec::new();
1665
1666    for i in 0..n {
1667        if i == idx {
1668            v.push(1.0);
1669        } else {
1670            v.push(0.0)
1671        }
1672    }
1673
1674    return v;
1675}
1676
1677impl Evaluator for MLEvaluator {
1678    fn eval_func(&self, b: &Board) -> i32 {
1679        let (att, def) = b.get_att_def();
1680
1681        let mut att_vec = Vec::new();
1682        let mut def_vec = Vec::new();
1683        for i in 0..64 {
1684            if (att >> i) & 1 == 1 {
1685                att_vec.push(1.0);
1686            } else {
1687                att_vec.push(0.0);
1688            }
1689
1690            if (def >> i) & 1 == 1 {
1691                def_vec.push(1.0);
1692            } else {
1693                def_vec.push(0.0);
1694            }
1695        }
1696
1697        let onehot = [att_vec, def_vec].concat();
1698        let onehot = Tensor::new(onehot, vec![128, 1]);
1699
1700        let val = self.g.inference(vec![onehot]);
1701
1702        return (val.get_item().unwrap() * MAX as f32) as i32;
1703    }
1704}
1705
1706impl GetAction for MLEvaluator {
1707    fn get_action(&self, b: &Board) -> u8 {
1708        let start = Instant::now();
1709        let (action, val, count) = self.eval_with_negalpha(b, 4);
1710        let end = start.elapsed();
1711        let time = end.as_nanos();
1712        println!(
1713            "[MLEvaluator]action:{action}, val:{val}, count:{count}, time:{}",
1714            time,
1715        );
1716        return action;
1717    }
1718}
1719
1720pub fn b2u128(b: &Board) -> u128 {
1721    let (att, def) = b.get_att_def();
1722    return (att as u128) | ((def as u128) << 64);
1723}
1724
1725pub fn u128_to_b(b: u128) -> Board {
1726    let mut board = Board::new();
1727    let black = b as u64;
1728    let white = (b >> 64) as u64;
1729    board.black = black;
1730    board.white = white;
1731
1732    return board;
1733}
1734
1735pub struct NNUE {
1736    pub g: Graph,
1737    loss: usize,
1738    g_out: usize,
1739    input: usize,
1740    t: usize,
1741    pub w1: usize,
1742    pub w1_size: usize,
1743    base_vec: Vec<Vec<f32>>,
1744    depth: usize,
1745}
1746
1747impl NNUE {
1748    pub fn new(g: Graph) -> Self {
1749        return NNUE {
1750            g: g,
1751            loss: 0,
1752            g_out: 0,
1753            input: 0,
1754            t: 0,
1755            w1: 0,
1756            w1_size: 0,
1757            base_vec: Vec::new(),
1758            depth: 0,
1759        };
1760    }
1761
1762    pub fn default() -> Self {
1763        use super::ml::*;
1764        use super::ml::{funcs::*, optim::*, params::*};
1765
1766        let w1_size = 256;
1767        let middle_size = 32;
1768
1769        let mut g = Graph::new();
1770        g.optimizer = Some(Box::new(MomentumSGD::new(0.01, 0.9)));
1771        let i1: usize = g.push_placeholder();
1772        let i2: usize = g.push_placeholder();
1773        // let t = g.push_placeholder();
1774
1775        let w1 = MM::auto(128, w1_size);
1776        let w1 = g.add_layer(vec![i1], Box::new(w1));
1777
1778        let l1 = Bias::auto(w1_size);
1779        let l1 = g.add_layer(vec![w1], Box::new(l1));
1780
1781        let activate = LeaklyReLU::default();
1782        let relu = g.add_layer(vec![l1], Box::new(activate));
1783
1784        let l2 = Linear::auto(w1_size, middle_size);
1785        let l2 = g.add_layer(vec![relu], Box::new(l2));
1786        let relu2 = g.add_layer(vec![l2], Box::new(LeaklyReLU::default()));
1787
1788        let l3 = Linear::auto(middle_size, middle_size);
1789        let l3 = g.add_layer(vec![relu2], Box::new(l3));
1790        let relu3 = g.add_layer(vec![l3], Box::new(LeaklyReLU::default()));
1791
1792        let l4 = Linear::auto(middle_size, 1);
1793        let l4 = g.add_layer(vec![relu3], Box::new(l4));
1794
1795        // t = lambda * result + (1 - lambda) * t_in
1796        let sig = g.add_layer(vec![l4], Box::new(Sigmoid::new(1.0)));
1797        let loss = g.add_layer(vec![sig, i2], Box::new(BinaryCrossEntropy::default()));
1798        // let loss = g.add_layer(vec![sig, i2], Box::new(MSE::new()));
1799
1800        g.set_target(sig);
1801        g.set_placeholder(vec![i1]);
1802
1803        return NNUE {
1804            g: g,
1805            loss: loss,
1806            g_out: sig,
1807            input: i1,
1808            t: i2,
1809            w1: w1,
1810            w1_size: w1_size,
1811            base_vec: Vec::new(),
1812            depth: 3,
1813        };
1814    }
1815
1816    pub fn inference(&self, b: &Board) -> f32 {
1817        let onehot = u2vec(Self::b2u128(b));
1818        let onehot = Tensor::new(onehot, vec![128, 1]);
1819
1820        let val = self.g.inference(vec![onehot]);
1821
1822        return val.get_item().unwrap();
1823    }
1824
1825    fn b2u128(b: &Board) -> u128 {
1826        let (att, def) = b.get_att_def();
1827        return (att as u128) | ((def as u128) << 64);
1828    }
1829
1830    pub fn set_inference(&mut self) {
1831        self.set_before_w1();
1832        for i in 0..128 {
1833            let onehot = onehot_vec(128, i);
1834            let onehot = Tensor::new(onehot, vec![128, 1]);
1835
1836            let val = self.g.inference(vec![onehot]);
1837            self.base_vec.push(val.data);
1838        }
1839        // println!("{:?}", self.base_vec);
1840        self.set_after_w1();
1841    }
1842
1843    pub fn set_depth(&mut self, depth: usize) {
1844        self.depth = depth;
1845    }
1846
1847    pub fn eval_with_negalpha(&self, b: &Board) -> (u8, f32, i32) {
1848        let b_hash = Self::b2u128(b);
1849        let b_vec = self.create_diff_vec(0, b_hash);
1850        // let start = Instant::now();
1851        // let result = eval_actor(&m7, &m6, 10, false);
1852        // let (a, b, c) = self.eval_with_negalpha_1(b, b_hash, b_vec, self.depth as u8, -2.0, 2.0);
1853        let (a, b, c) =
1854            self.eval_with_negalpha_(b.clone(), b_hash, b_vec, None, self.depth as u8, -2.0, 2.0);
1855        return (a, b, c);
1856    }
1857
1858    pub fn eval_with_negalpha_(
1859        &self,
1860        b: Board,
1861        b_hash: u128,
1862        b_vec: Vec<f32>,
1863        next: Option<(u128, &Vec<f32>)>,
1864        depth: u8,
1865        alpha: f32,
1866        beta: f32,
1867    ) -> (u8, f32, i32) {
1868        use std::cmp::Ordering::*;
1869
1870        let mut count = 0;
1871        let actions = b.valid_actions();
1872        let mut max_val = -2.0;
1873        let mut max_action: u8 = 16;
1874        let mut next_info: Option<(u128, &Vec<f32>)> = next;
1875        let mut alpha = alpha;
1876
1877        if depth <= 1 {
1878            let mut next_vec;
1879            for &action in actions.iter() {
1880                let next_board = &b.next(action);
1881
1882                let next_hash = Self::b2u128(next_board);
1883
1884                if next_board.is_win() {
1885                    return (action, 1.0, count);
1886                } else if next_board.is_draw() {
1887                    return (action, 0.5, count);
1888                }
1889
1890                let feed_vec: Vec<f32>;
1891                match next_info {
1892                    None => {
1893                        next_vec = self.create_diff_vec(0, next_hash);
1894                        next_info = Some((next_hash, &next_vec));
1895                        feed_vec = next_vec.clone();
1896                    }
1897                    Some((hash, ref vec)) => {
1898                        let next_vec_ = self.create_diff_vec(hash, next_hash);
1899
1900                        feed_vec = next_vec_
1901                            .iter()
1902                            .zip(vec.iter())
1903                            .map(|(a, b)| a + b)
1904                            .collect();
1905                    }
1906                }
1907
1908                let w1 = Tensor::new(feed_vec, vec![self.w1_size, 1]);
1909
1910                let val = 1.0 - self.g.inference(vec![w1]).get_item().unwrap();
1911
1912                count += 1;
1913                if max_val < val {
1914                    max_val = val;
1915                    max_action = action;
1916                    if max_val > alpha {
1917                        alpha = max_val;
1918                        if alpha > beta {
1919                            return (max_action, max_val, count);
1920                        }
1921                    }
1922                }
1923            }
1924        } else {
1925            let mut nexts = Vec::new();
1926            let mut next_vec;
1927            for &action in actions.iter() {
1928                let next_board = b.next(action);
1929                let next_hash = Self::b2u128(&next_board);
1930
1931                if next_board.is_win() {
1932                    return (action, 1.0, count);
1933                } else if next_board.is_draw() {
1934                    return (action, 0.5, count);
1935                }
1936
1937                let feed_vec: Vec<f32>;
1938                match next_info {
1939                    None => {
1940                        next_vec = self.create_diff_vec(0, next_hash);
1941                        next_info = Some((next_hash, &next_vec));
1942                        feed_vec = next_vec.clone();
1943                    }
1944                    Some((hash, ref vec)) => {
1945                        let next_vec_ = self.create_diff_vec(hash, next_hash);
1946                        feed_vec = next_vec_
1947                            .iter()
1948                            .zip(vec.iter())
1949                            .map(|(a, b)| a + b)
1950                            .collect();
1951                    }
1952                }
1953                let w1 = Tensor::new(feed_vec.clone(), vec![self.w1_size, 1]);
1954
1955                let val = 1.0 - self.g.inference(vec![w1]).get_item().unwrap();
1956
1957                nexts.push((action, next_board, next_hash, feed_vec, val))
1958            }
1959            nexts.sort_by(|a, b| {
1960                if a.4 < b.4 {
1961                    return Greater;
1962                } else {
1963                    return Less;
1964                }
1965            });
1966
1967            for (action, next_board, next_hash, next_vec, val) in nexts {
1968                let (_, val, _count) = self.eval_with_negalpha_(
1969                    next_board,
1970                    next_hash,
1971                    next_vec,
1972                    Some((b_hash, &b_vec)),
1973                    depth - 1,
1974                    1.0 - beta,
1975                    1.0 - alpha,
1976                );
1977                let val = -0.999 * (val - 0.5) + 0.5;
1978
1979                count += _count;
1980                if max_val < val {
1981                    max_val = val;
1982                    max_action = action;
1983                    if max_val > alpha {
1984                        alpha = max_val;
1985                        if alpha > beta {
1986                            return (max_action, max_val, count);
1987                        }
1988                    }
1989                }
1990            }
1991        }
1992        return (max_action, max_val, count);
1993    }
1994
1995    fn create_diff_vec(&self, a: u128, b: u128) -> Vec<f32> {
1996        // a -> b を考える
1997        let minus = a & !b;
1998        let plus = b & !a;
1999
2000        let mut minus_vec = vec![0.0; self.w1_size];
2001        let mut plus_vec = vec![0.0; self.w1_size];
2002
2003        for i in 0..128 {
2004            if (minus >> i) & 1 == 1 {
2005                for j in 0..self.w1_size {
2006                    minus_vec[j] += self.base_vec[i][j];
2007                }
2008                continue;
2009            }
2010            if (plus >> i) & 1 == 1 {
2011                for j in 0..self.w1_size {
2012                    plus_vec[j] += self.base_vec[i][j];
2013                }
2014            }
2015        }
2016
2017        for i in 0..self.w1_size {
2018            plus_vec[i] -= minus_vec[i];
2019        }
2020
2021        return plus_vec;
2022    }
2023
2024    pub fn train(&mut self) {
2025        self.g.set_placeholder(vec![self.input, self.t]);
2026        self.g.set_target(self.loss);
2027    }
2028
2029    pub fn eval(&mut self) {
2030        self.g.set_placeholder(vec![self.input]);
2031        self.g.set_target(self.g_out);
2032    }
2033
2034    pub fn set_before_w1(&mut self) {
2035        self.g.set_placeholder(vec![self.input]);
2036        self.g.set_target(self.w1);
2037    }
2038
2039    pub fn set_after_w1(&mut self) {
2040        self.g.set_placeholder(vec![self.w1]);
2041        self.g.set_target(self.g_out)
2042    }
2043
2044    pub fn save(&self, s: String) {
2045        self.g.save(s);
2046    }
2047
2048    pub fn load(&mut self, s: String) {
2049        self.g.load(s);
2050    }
2051}
2052
2053impl GetAction for NNUE {
2054    fn get_action(&self, b: &Board) -> u8 {
2055        let start = Instant::now();
2056        let (action, val, count) = self.eval_with_negalpha(b);
2057        let end = start.elapsed();
2058        let hoge: Box<dyn Evaluator> = Box::new(CoEvaluator::best());
2059        let (action_, val_, count) = negalpha(b, 3, -MAX - 1, MAX + 1, &hoge);
2060        let time = end.as_nanos();
2061        if cfg!(feature = "render") {
2062            println!(
2063                "[NNUE]action:{action}-{action_}, val:{val}, val_:{val_}, count:{count}, time:{}",
2064                time,
2065            );
2066        }
2067        let res = mate_check_horizontal(b);
2068        if let Some((flag, action)) = res {
2069            if cfg!(feature = "render") {
2070                println!("{flag}");
2071            }
2072            return action;
2073        }
2074        return action;
2075    }
2076}
2077
2078impl Analyzer for NNUE {
2079    fn analyze_eval(&self, b: &Board) -> f32 {
2080        let (action, val, count) = self.eval_with_negalpha(b);
2081        return val;
2082    }
2083}
2084
2085pub struct RowEvaluator {
2086    w_row3: i32,
2087    w_row2: i32,
2088    w_row1: i32,
2089}
2090
2091impl RowEvaluator {
2092    pub fn new() -> Self {
2093        return RowEvaluator {
2094            w_row1: 1,
2095            w_row2: 1,
2096            w_row3: 1,
2097        };
2098    }
2099
2100    pub fn best() -> Self {
2101        RowEvaluator::from(2, 5, 12)
2102    }
2103
2104    pub fn from(w1: i32, w2: i32, w3: i32) -> Self {
2105        return RowEvaluator {
2106            w_row3: w3,
2107            w_row2: w2,
2108            w_row1: w1,
2109        };
2110    }
2111}
2112
2113impl Evaluator for RowEvaluator {
2114    fn eval_func(&self, b: &Board) -> i32 {
2115        let (att, def) = b.get_att_def();
2116        let blank = !att & !def;
2117
2118        let att3 = count_3row(att, blank) as i32;
2119        let def3 = count_3row(def, blank) as i32;
2120
2121        let att2 = count_2row(att, blank) as i32;
2122        let def2 = count_2row(def, blank) as i32;
2123
2124        let att1 = count_1row(att, blank) as i32;
2125        let def1 = count_1row(def, blank) as i32;
2126
2127        // pprint_board(b);
2128        // println!("{att}, {def}");
2129
2130        return self.w_row3 * (att3 - def3)
2131            + self.w_row2 * (att2 - def2)
2132            + self.w_row1 * (att1 - def1);
2133    }
2134}
2135
2136pub struct NullEvaluator {}
2137impl NullEvaluator {
2138    pub fn new() -> Self {
2139        return NullEvaluator {};
2140    }
2141}
2142
2143impl Evaluator for NullEvaluator {
2144    fn eval_func(&self, b: &Board) -> i32 {
2145        return 0;
2146    }
2147}
2148
2149pub struct CoEvaluator {
2150    a: Box<dyn Evaluator>,
2151    b: Box<dyn Evaluator>,
2152    a_weight: i32,
2153    b_weight: i32,
2154}
2155
2156impl CoEvaluator {
2157    pub fn new(a: Box<dyn Evaluator>, b: Box<dyn Evaluator>, a_weight: i32, b_weight: i32) -> Self {
2158        return CoEvaluator {
2159            a: a,
2160            b: b,
2161            a_weight: a_weight,
2162            b_weight: b_weight,
2163        };
2164    }
2165
2166    pub fn best() -> Self {
2167        let a = RowEvaluator::best();
2168        let b = PositionEvaluator::best();
2169        return CoEvaluator {
2170            a: Box::new(a),
2171            b: Box::new(b),
2172            a_weight: 3,
2173            b_weight: 1,
2174        };
2175    }
2176}
2177
2178impl Evaluator for CoEvaluator {
2179    fn eval_func(&self, b: &Board) -> i32 {
2180        let a_score = self.a.eval_func(b);
2181        let b_score = self.b.eval_func(b);
2182
2183        let score = 10 * (a_score * self.a_weight + b_score * self.b_weight)
2184            / (self.a_weight + self.b_weight);
2185
2186        return score;
2187    }
2188}
2189
2190impl EvaluatorF for CoEvaluator {
2191    fn eval_func_f32(&self, b: &Board) -> f32 {
2192        let val = self.eval_func(b) as f32;
2193        return 1.0 / (1.0 + (-val / 400.0).exp());
2194    }
2195}
2196
2197#[derive(Clone)]
2198pub struct LineEvaluator {
2199    pub w_float_3_line: [f32; 32],
2200    pub w_ground_3_line: [f32; 32],
2201    pub w_float_2_line: [f32; 64],
2202    pub w_ground_2_line: [f32; 64],
2203    pub w_float_1_line: [f32; 96],
2204    pub w_ground_1_line: [f32; 96],
2205    pub w_dif_float_3_line: [f32; 64],
2206    pub w_dif_ground_3_line: [f32; 64],
2207    pub w_dif_float_2_line: [f32; 128],
2208    pub w_dif_ground_2_line: [f32; 128],
2209    pub w_dif_float_1_line: [f32; 192],
2210    pub w_dif_ground_1_line: [f32; 192],
2211    pub w_pos_3_line: [f32; 64],
2212    pub w_pos_2_line: [f32; 64],
2213    pub w_pos_1_line: [f32; 64],
2214    pub w_dpos_3_line: [f32; 64],
2215    pub w_dpos_2_line: [f32; 64],
2216    pub w_dpos_1_line: [f32; 64],
2217    pub w_bf_3_line: [f32; 64],
2218    pub w_bf_2_line: [f32; 64],
2219    pub w_bf_1_line: [f32; 64],
2220    pub w_bg_3_line: [f32; 64],
2221    pub w_bg_2_line: [f32; 64],
2222    pub w_bg_1_line: [f32; 64],
2223    pub w_dbf_3_line: [f32; 64],
2224    pub w_dbf_2_line: [f32; 64],
2225    pub w_dbf_1_line: [f32; 64],
2226    pub w_dbg_3_line: [f32; 64],
2227    pub w_dbg_2_line: [f32; 64],
2228    pub w_dbg_1_line: [f32; 64],
2229    pub w_pos: [f32; 64],
2230    pub w_dpos: [f32; 64],
2231    pub w_trap_3_num: [f32; 13],
2232    pub bias: f32,
2233}
2234
2235type LineMaskBundle = (
2236    u64,
2237    u64,
2238    u64,
2239    u64,
2240    u64,
2241    u64,
2242    u64,
2243    u64,
2244    u64,
2245    u64,
2246    u64,
2247    u64,
2248    u64,
2249);
2250
2251pub fn acum_or(bundle: LineMaskBundle) -> u64 {
2252    return bundle.0
2253        | bundle.1
2254        | bundle.2
2255        | bundle.3
2256        | bundle.4
2257        | bundle.5
2258        | bundle.6
2259        | bundle.7
2260        | bundle.8
2261        | bundle.9
2262        | bundle.10
2263        | bundle.11
2264        | bundle.12;
2265}
2266
2267pub fn acum_mask_bundle(bundle: LineMaskBundle) -> u32 {
2268    return bundle.0.count_ones()
2269        + bundle.1.count_ones()
2270        + bundle.2.count_ones()
2271        + bundle.3.count_ones()
2272        + bundle.4.count_ones()
2273        + bundle.5.count_ones()
2274        + bundle.6.count_ones()
2275        + bundle.7.count_ones()
2276        + bundle.8.count_ones()
2277        + bundle.9.count_ones()
2278        + bundle.10.count_ones()
2279        + bundle.11.count_ones()
2280        + bundle.12.count_ones();
2281}
2282pub fn apply_mask_bundle(bundle: LineMaskBundle, mask: u64) -> LineMaskBundle {
2283    return (
2284        bundle.0 & mask,
2285        bundle.1 & mask,
2286        bundle.2 & mask,
2287        bundle.3 & mask,
2288        bundle.4 & mask,
2289        bundle.5 & mask,
2290        bundle.6 & mask,
2291        bundle.7 & mask,
2292        bundle.8 & mask,
2293        bundle.9 & mask,
2294        bundle.10 & mask,
2295        bundle.11 & mask,
2296        bundle.12 & mask,
2297    );
2298}
2299
2300impl LineEvaluator {
2301    pub fn new() -> Self {
2302        return LineEvaluator {
2303            w_float_3_line: [0.0; 32],
2304            w_ground_3_line: [0.0; 32],
2305            w_float_2_line: [0.0; 64],
2306            w_ground_2_line: [0.0; 64],
2307            w_float_1_line: [0.0; 96],
2308            w_ground_1_line: [0.0; 96],
2309            w_dif_float_3_line: [0.0; 64],
2310            w_dif_ground_3_line: [0.0; 64],
2311            w_dif_float_2_line: [0.0; 128],
2312            w_dif_ground_2_line: [0.0; 128],
2313            w_dif_float_1_line: [0.0; 192],
2314            w_dif_ground_1_line: [0.0; 192],
2315            w_pos_3_line: [0.0; 64],
2316            w_pos_2_line: [0.0; 64],
2317            w_pos_1_line: [0.0; 64],
2318            w_dpos_3_line: [0.0; 64],
2319            w_dpos_2_line: [0.0; 64],
2320            w_dpos_1_line: [0.0; 64],
2321            w_bf_3_line: [0.0; 64],
2322            w_bg_3_line: [0.0; 64],
2323            w_bf_2_line: [0.0; 64],
2324            w_bg_2_line: [0.0; 64],
2325            w_bf_1_line: [0.0; 64],
2326            w_bg_1_line: [0.0; 64],
2327            w_dbf_3_line: [0.0; 64],
2328            w_dbg_3_line: [0.0; 64],
2329            w_dbf_2_line: [0.0; 64],
2330            w_dbg_2_line: [0.0; 64],
2331            w_dbf_1_line: [0.0; 64],
2332            w_dbg_1_line: [0.0; 64],
2333            w_pos: [0.0; 64],
2334            w_dpos: [0.0; 64],
2335            w_trap_3_num: [0.0; 13],
2336            bias: 0.0,
2337        };
2338    }
2339
2340    pub fn analyze_line(
2341        a1: u64,
2342        a2: u64,
2343        a3: u64,
2344        a4: u64,
2345        b1: u64,
2346        b2: u64,
2347        b3: u64,
2348        b4: u64,
2349        mask: u64,
2350        magic: u64,
2351    ) -> (u64, u64, u64) {
2352        return (
2353            ((b1 & b2 & b3 & a4 | b1 & b2 & a3 & b4 | b1 & a2 & b3 & b4 | a1 & b2 & b3 & b4)
2354                & mask)
2355                * magic,
2356            ((a1 & a2 & b3 & b4
2357                | a1 & b2 & a3 & b4
2358                | a1 & b2 & b3 & a4
2359                | b1 & a2 & a3 & b4
2360                | b1 & a2 & b3 & a4
2361                | b1 & b2 & a3 & a4)
2362                & mask)
2363                * magic,
2364            ((a1 & a2 & a3 & b4 | a1 & a2 & b3 & a4 | a1 & b2 & a3 & a4 | b1 & a2 & a3 & a4)
2365                & mask)
2366                * magic,
2367        );
2368    }
2369
2370    pub fn analyze_board(
2371        a: u64,
2372        d: u64,
2373    ) -> (
2374        LineMaskBundle,
2375        LineMaskBundle,
2376        LineMaskBundle,
2377        LineMaskBundle,
2378        LineMaskBundle,
2379        LineMaskBundle,
2380    ) {
2381        let stone = a | d;
2382        let b = !stone;
2383        let (
2384            a1,
2385            a2,
2386            a3,
2387            a4,
2388            a5,
2389            a6,
2390            a8,
2391            a9,
2392            a10,
2393            a11,
2394            a12,
2395            a13,
2396            a15,
2397            a16,
2398            a17,
2399            a19,
2400            a20,
2401            a21,
2402            a22,
2403            a24,
2404            a26,
2405            a30,
2406            a32,
2407            a33,
2408            a34,
2409            a36,
2410            a38,
2411            a39,
2412            a40,
2413            a42,
2414            a45,
2415            a48,
2416            a51,
2417            a57,
2418            a60,
2419            a63,
2420        ) = (
2421            a >> 1,
2422            a >> 2,
2423            a >> 3,
2424            a >> 4,
2425            a >> 5,
2426            a >> 6,
2427            a >> 8,
2428            a >> 9,
2429            a >> 10,
2430            a >> 11,
2431            a >> 12,
2432            a >> 13,
2433            a >> 15,
2434            a >> 16,
2435            a >> 17,
2436            a >> 19,
2437            a >> 20,
2438            a >> 21,
2439            a >> 22,
2440            a >> 24,
2441            a >> 26,
2442            a >> 30,
2443            a >> 32,
2444            a >> 33,
2445            a >> 34,
2446            a >> 36,
2447            a >> 38,
2448            a >> 39,
2449            a >> 40,
2450            a >> 42,
2451            a >> 45,
2452            a >> 48,
2453            a >> 51,
2454            a >> 57,
2455            a >> 60,
2456            a >> 63,
2457        );
2458        let (
2459            b1,
2460            b2,
2461            b3,
2462            b4,
2463            b5,
2464            b6,
2465            b8,
2466            b9,
2467            b10,
2468            b11,
2469            b12,
2470            b13,
2471            b15,
2472            b16,
2473            b17,
2474            b19,
2475            b20,
2476            b21,
2477            b22,
2478            b24,
2479            b26,
2480            b30,
2481            b32,
2482            b33,
2483            b34,
2484            b36,
2485            b38,
2486            b39,
2487            b40,
2488            b42,
2489            b45,
2490            b48,
2491            b51,
2492            b57,
2493            b60,
2494            b63,
2495        ) = (
2496            b >> 1,
2497            b >> 2,
2498            b >> 3,
2499            b >> 4,
2500            b >> 5,
2501            b >> 6,
2502            b >> 8,
2503            b >> 9,
2504            b >> 10,
2505            b >> 11,
2506            b >> 12,
2507            b >> 13,
2508            b >> 15,
2509            b >> 16,
2510            b >> 17,
2511            b >> 19,
2512            b >> 20,
2513            b >> 21,
2514            b >> 22,
2515            b >> 24,
2516            b >> 26,
2517            b >> 30,
2518            b >> 32,
2519            b >> 33,
2520            b >> 34,
2521            b >> 36,
2522            b >> 38,
2523            b >> 39,
2524            b >> 40,
2525            b >> 42,
2526            b >> 45,
2527            b >> 48,
2528            b >> 51,
2529            b >> 57,
2530            b >> 60,
2531            b >> 63,
2532        );
2533        let (x1, x2, x3) =
2534            LineEvaluator::analyze_line(a, a1, a2, a3, b, b1, b2, b3, 0x1111_1111_1111_1111, 0xf);
2535        let (x1, x2, x3, px1, px2, px3) = (x1 & b, x2 & b, x3 & b, x1 & a, x2 & a, x3 & a);
2536
2537        let (y1, y2, y3) = LineEvaluator::analyze_line(
2538            a,
2539            a4,
2540            a8,
2541            a12,
2542            b,
2543            b4,
2544            b8,
2545            b12,
2546            0x000f_000f_000f_000f,
2547            0x1111,
2548        );
2549        let (y1, y2, y3, py1, py2, py3) = (y1 & b, y2 & b, y3 & b, y1 & a, y2 & a, y3 & a);
2550
2551        let (z1, z2, z3) = LineEvaluator::analyze_line(
2552            a,
2553            a16,
2554            a32,
2555            a48,
2556            b,
2557            b16,
2558            b32,
2559            b48,
2560            0xffff,
2561            0x0001_0001_0001_0001,
2562        );
2563        let (xy1, xy2, xy3) = LineEvaluator::analyze_line(
2564            a,
2565            a5,
2566            a10,
2567            a15,
2568            b,
2569            b5,
2570            b10,
2571            b15,
2572            0x0001_0001_0001_0001,
2573            0x8421,
2574        );
2575        let (xy1_, xy2_, xy3_) =
2576            LineEvaluator::analyze_line(a, a3, a6, a9, b, b3, b6, b9, 0x0008_0008_0008_0008, 0x249);
2577        let (xz1, xz2, xz3) = LineEvaluator::analyze_line(
2578            a,
2579            a17,
2580            a34,
2581            a51,
2582            b,
2583            b17,
2584            b34,
2585            b51,
2586            0x1111,
2587            0x0008_0004_0002_0001,
2588        );
2589        let (xz1_, xz2_, xz3_) = LineEvaluator::analyze_line(
2590            a,
2591            a15,
2592            a30,
2593            a45,
2594            b,
2595            b15,
2596            b30,
2597            b45,
2598            0x8888,
2599            0x2000_4000_8001,
2600        );
2601        let (yz1, yz2, yz3) = LineEvaluator::analyze_line(
2602            a,
2603            a20,
2604            a40,
2605            a60,
2606            b,
2607            b20,
2608            b40,
2609            b60,
2610            0x000f,
2611            0x1000_0100_0010_0001,
2612        );
2613        let (yz1_, yz2_, yz3_) = LineEvaluator::analyze_line(
2614            a,
2615            a12,
2616            a24,
2617            a36,
2618            b,
2619            b12,
2620            b24,
2621            b36,
2622            0xf000,
2623            0x0000_0010_0100_1001,
2624        );
2625        let (xyz11, xyz12, xyz13) = LineEvaluator::analyze_line(
2626            a,
2627            a21,
2628            a42,
2629            a63,
2630            b,
2631            b21,
2632            b42,
2633            b63,
2634            0x1,
2635            0x8000_0400_0020_0001,
2636        );
2637        let (xyz21, xyz22, xyz23) = LineEvaluator::analyze_line(
2638            a,
2639            a19,
2640            a38,
2641            a57,
2642            b,
2643            b19,
2644            b38,
2645            b57,
2646            0x8,
2647            0x0200_0040_0008_0001,
2648        );
2649        let (xyz31, xyz32, xyz33) = LineEvaluator::analyze_line(
2650            a,
2651            a13,
2652            a26,
2653            a39,
2654            b,
2655            b13,
2656            b26,
2657            b39,
2658            0x1000,
2659            0x0000_0080_0400_2001,
2660        );
2661        let (xyz41, xyz42, xyz43) = LineEvaluator::analyze_line(
2662            a,
2663            a11,
2664            a22,
2665            a33,
2666            b,
2667            b11,
2668            b22,
2669            b33,
2670            0x8000,
2671            0x0000_0002_0040_0801,
2672        );
2673
2674        let (z1, z2, z3, pz1, pz2, pz3) = (z1 & b, z2 & b, z3 & b, z1 & a, z2 & a, z3 & a);
2675        let (xy1, xy2, xy3, pxy1, pxy2, pxy3) =
2676            (xy1 & b, xy2 & b, xy3 & b, xy1 & a, xy2 & a, xy3 & a);
2677        let (xy1_, xy2_, xy3_, pxy1_, pxy2_, pxy3_) =
2678            (xy1_ & b, xy2_ & b, xy3_ & b, xy1_ & a, xy2_ & a, xy3_ & a);
2679        let (yz1, yz2, yz3, pyz1, pyz2, pyz3) =
2680            (yz1 & b, yz2 & b, yz3 & b, yz1 & a, yz2 & a, yz3 & a);
2681        let (yz1_, yz2_, yz3_, pyz1_, pyz2_, pyz3_) =
2682            (yz1_ & b, yz2_ & b, yz3_ & b, yz1_ & a, yz2_ & a, yz3_ & a);
2683        let (xz1, xz2, xz3, pxz1, pxz2, pxz3) =
2684            (xz1 & b, xz2 & b, xz3 & b, xz1 & a, xz2 & a, xz3 & a);
2685        let (xz1_, xz2_, xz3_, pxz1_, pxz2_, pxz3_) =
2686            (xz1_ & b, xz2_ & b, xz3_ & b, xz1_ & a, xz2_ & a, xz3_ & a);
2687        let (xyz11, xyz12, xyz13, pxyz11, pxyz12, pxyz13) = (
2688            xyz11 & b,
2689            xyz12 & b,
2690            xyz13 & b,
2691            xyz11 & a,
2692            xyz12 & a,
2693            xyz13 & a,
2694        );
2695        let (xyz21, xyz22, xyz23, pxyz21, pxyz22, pxyz23) = (
2696            xyz21 & b,
2697            xyz22 & b,
2698            xyz23 & b,
2699            xyz21 & a,
2700            xyz22 & a,
2701            xyz23 & a,
2702        );
2703        let (xyz31, xyz32, xyz33, pxyz31, pxyz32, pxyz33) = (
2704            xyz31 & b,
2705            xyz32 & b,
2706            xyz33 & b,
2707            xyz31 & a,
2708            xyz32 & a,
2709            xyz33 & a,
2710        );
2711        let (xyz41, xyz42, xyz43, pxyz41, pxyz42, pxyz43) = (
2712            xyz41 & b,
2713            xyz42 & b,
2714            xyz43 & b,
2715            xyz41 & a,
2716            xyz42 & a,
2717            xyz43 & a,
2718        );
2719
2720        return (
2721            (
2722                x1, y1, z1, xy1, xy1_, yz1, yz1_, xz1, xz1_, xyz11, xyz21, xyz31, xyz41,
2723            ),
2724            (
2725                x2, y2, z2, xy2, xy2_, yz2, yz2_, xz2, xz2_, xyz12, xyz22, xyz32, xyz42,
2726            ),
2727            (
2728                x3, y3, z3, xy3, xy3_, yz3, yz3_, xz3, xz3_, xyz13, xyz23, xyz33, xyz43,
2729            ),
2730            (
2731                px1, py1, pz1, pxy1, pxy1_, pyz1, pyz1_, pxz1, pxz1_, pxyz11, pxyz21, pxyz31,
2732                pxyz41,
2733            ),
2734            (
2735                px2, py2, pz2, pxy2, pxy2_, pyz2, pyz2_, pxz2, pxz2_, pxyz12, pxyz22, pxyz32,
2736                pxyz42,
2737            ),
2738            (
2739                px3, py3, pz3, pxy3, pxy3_, pyz3, pyz3_, pxz3, pxz3_, pxyz13, pxyz23, pxyz33,
2740                pxyz43,
2741            ),
2742        );
2743    }
2744
2745    pub fn get_counts(
2746        &self,
2747        b: &Board,
2748    ) -> (
2749        usize,
2750        usize,
2751        usize,
2752        usize,
2753        usize,
2754        usize,
2755        usize,
2756        usize,
2757        usize,
2758        usize,
2759        usize,
2760        usize,
2761        u64,
2762        u64,
2763        u64,
2764        u64,
2765        u64,
2766        u64,
2767        u64,
2768        u64,
2769        u64,
2770        u64,
2771        u64,
2772        u64,
2773        u64,
2774        u64,
2775        u64,
2776        u64,
2777        u64,
2778        u64,
2779        usize,
2780        u64,
2781        u64,
2782    ) {
2783        let (att, def) = b.get_att_def();
2784        let (a1, a2, a3, pa1, pa2, pa3) = LineEvaluator::analyze_board(att, def);
2785        let (d1, d2, d3, pd1, pd2, pd3) = LineEvaluator::analyze_board(def, att);
2786        let stone = att | def;
2787        let ground = !stone & (stone << 16 | 0xffff);
2788        let float = !stone ^ ground;
2789        let a1_float = acum_mask_bundle(apply_mask_bundle(a1, float)) as usize;
2790        let a1_ground = acum_mask_bundle(apply_mask_bundle(a1, ground)) as usize;
2791        let a2_float = acum_mask_bundle(apply_mask_bundle(a2, float)) as usize;
2792        let a2_ground = acum_mask_bundle(apply_mask_bundle(a2, ground)) as usize;
2793        let a3_float = acum_mask_bundle(apply_mask_bundle(a3, float)) as usize;
2794        let a3_ground = acum_mask_bundle(apply_mask_bundle(a3, ground)) as usize;
2795        let d1_float = acum_mask_bundle(apply_mask_bundle(d1, float)) as usize;
2796        let d1_ground = acum_mask_bundle(apply_mask_bundle(d1, ground)) as usize;
2797        let d2_float = acum_mask_bundle(apply_mask_bundle(d2, float)) as usize;
2798        let d2_ground = acum_mask_bundle(apply_mask_bundle(d2, ground)) as usize;
2799        let d3_float = acum_mask_bundle(apply_mask_bundle(d3, float)) as usize;
2800        let d3_ground = acum_mask_bundle(apply_mask_bundle(d3, ground)) as usize;
2801
2802        let af1_mask = acum_or(a1) & float;
2803        let ag1_mask = acum_or(a1) & ground;
2804        let af2_mask = acum_or(a2) & float;
2805        let ag2_mask = acum_or(a2) & ground;
2806        let af3_mask = acum_or(a3) & float;
2807        let ag3_mask = acum_or(a3) & ground;
2808        let df1_mask = acum_or(d1) & float;
2809        let dg1_mask = acum_or(d1) & ground;
2810        let df2_mask = acum_or(d2) & float;
2811        let dg2_mask = acum_or(d2) & ground;
2812        let df3_mask = acum_or(d3) & float;
2813        let dg3_mask = acum_or(d3) & ground;
2814        let trap_3_num = (((acum_or(d3) | acum_or(a3)) & 0x0000_ffff_0000_0000).count_ones()
2815            + (stone.count_ones() % 2)) as usize;
2816
2817        let pa1_mask = acum_or(pa1);
2818        let pa2_mask = acum_or(pa2);
2819        let pa3_mask = acum_or(pa3);
2820        let pd1_mask = acum_or(pd1);
2821        let pd2_mask = acum_or(pd2);
2822        let pd3_mask = acum_or(pd3);
2823
2824        return (
2825            a1_float, a2_float, a3_float, a1_ground, a2_ground, a3_ground, d1_float, d2_float,
2826            d3_float, d1_ground, d2_ground, d3_ground, af1_mask, af2_mask, af3_mask, ag1_mask,
2827            ag2_mask, ag3_mask, df1_mask, df2_mask, df3_mask, dg1_mask, dg2_mask, dg3_mask,
2828            pa1_mask, pa2_mask, pa3_mask, pd1_mask, pd2_mask, pd3_mask, trap_3_num, att, def,
2829        );
2830    }
2831
2832    pub fn evaluate_board(&self, b: &Board) -> f32 {
2833        let (att, def) = b.get_att_def();
2834        let (a1, a2, a3, pa1, pa2, pa3) = LineEvaluator::analyze_board(att, def);
2835        let (d1, d2, d3, pd1, pd2, pd3) = LineEvaluator::analyze_board(def, att);
2836        let stone = att | def;
2837        let ground = !stone & (stone << 16 | 0xffff);
2838        let float = !stone ^ ground;
2839        let a1_float = acum_mask_bundle(apply_mask_bundle(a1, float)) as usize;
2840        let a1_ground = acum_mask_bundle(apply_mask_bundle(a1, ground)) as usize;
2841        let a2_float = acum_mask_bundle(apply_mask_bundle(a2, float)) as usize;
2842        let a2_ground = acum_mask_bundle(apply_mask_bundle(a2, ground)) as usize;
2843        let a3_float = acum_mask_bundle(apply_mask_bundle(a3, float)) as usize;
2844        let a3_ground = acum_mask_bundle(apply_mask_bundle(a3, ground)) as usize;
2845        let d1_float = acum_mask_bundle(apply_mask_bundle(d1, float)) as usize;
2846        let d1_ground = acum_mask_bundle(apply_mask_bundle(d1, ground)) as usize;
2847        let d2_float = acum_mask_bundle(apply_mask_bundle(d2, float)) as usize;
2848        let d2_ground = acum_mask_bundle(apply_mask_bundle(d2, ground)) as usize;
2849        let d3_float = acum_mask_bundle(apply_mask_bundle(d3, float)) as usize;
2850        let d3_ground = acum_mask_bundle(apply_mask_bundle(d3, ground)) as usize;
2851
2852        let af1_mask = acum_or(a1) & float;
2853        let ag1_mask = acum_or(a1) & ground;
2854        let af2_mask = acum_or(a2) & float;
2855        let ag2_mask = acum_or(a2) & ground;
2856        let af3_mask = acum_or(a3) & float;
2857        let ag3_mask = acum_or(a3) & ground;
2858        let df1_mask = acum_or(d1) & float;
2859        let dg1_mask = acum_or(d1) & ground;
2860        let df2_mask = acum_or(d2) & float;
2861        let dg2_mask = acum_or(d2) & ground;
2862        let df3_mask = acum_or(d3) & float;
2863        let dg3_mask = acum_or(d3) & ground;
2864
2865        let pa1_mask = acum_or(pa1);
2866        let pa2_mask = acum_or(pa2);
2867        let pa3_mask = acum_or(pa3);
2868        let pd1_mask = acum_or(pd1);
2869        let pd2_mask = acum_or(pd2);
2870        let pd3_mask = acum_or(pd3);
2871
2872        let trap_3_num = (((acum_or(d3) | acum_or(a3)) & 0x0000_ffff_0000_0000).count_ones()
2873            + (stone.count_ones() % 2)) as usize;
2874
2875        let mut val = 0.0;
2876
2877        for i in 0..64 {
2878            let bit = 1 << i;
2879            if af1_mask & bit != 0 {
2880                val += self.w_bf_1_line[i];
2881            }
2882            if af2_mask & bit != 0 {
2883                val += self.w_bf_2_line[i];
2884            }
2885            if af3_mask & bit != 0 {
2886                val += self.w_bf_3_line[i];
2887            }
2888            if df1_mask & bit != 0 {
2889                val += self.w_dbf_1_line[i];
2890            }
2891            if df2_mask & bit != 0 {
2892                val += self.w_dbf_2_line[i];
2893            }
2894            if df3_mask & bit != 0 {
2895                val += self.w_dbf_3_line[i];
2896            }
2897            if ag1_mask & bit != 0 {
2898                val += self.w_bg_1_line[i];
2899            }
2900            if ag2_mask & bit != 0 {
2901                val += self.w_bg_2_line[i];
2902            }
2903            if ag3_mask & bit != 0 {
2904                val += self.w_bg_3_line[i];
2905            }
2906            if dg1_mask & bit != 0 {
2907                val += self.w_dbg_1_line[i];
2908            }
2909            if dg2_mask & bit != 0 {
2910                val += self.w_dbg_2_line[i];
2911            }
2912            if dg3_mask & bit != 0 {
2913                val += self.w_dbg_3_line[i];
2914            }
2915            if pa1_mask & bit != 0 {
2916                val += self.w_pos_1_line[i];
2917            } else if pd1_mask & bit != 0 {
2918                val += self.w_dpos_1_line[i];
2919            }
2920            if pa2_mask & bit != 0 {
2921                val += self.w_pos_2_line[i];
2922            } else if pd2_mask & bit != 0 {
2923                val += self.w_dpos_2_line[i];
2924            }
2925            if pa3_mask & bit != 0 {
2926                val += self.w_pos_3_line[i];
2927            } else if pd3_mask & bit != 0 {
2928                val += self.w_dpos_3_line[i];
2929            }
2930            if att & bit != 0 {
2931                val += self.w_pos[i];
2932            } else if def & bit != 0 {
2933                val += self.w_dpos[i];
2934            }
2935        }
2936
2937        val += self.w_float_1_line[a1_float]
2938            + self.w_float_2_line[a2_float]
2939            + self.w_float_3_line[a3_float]
2940            + self.w_ground_1_line[a1_ground]
2941            + self.w_ground_2_line[a2_ground]
2942            + self.w_ground_3_line[a3_ground]
2943            + (self.w_dif_float_1_line[96 - d1_float + a1_float]
2944                + self.w_dif_float_2_line[64 - d2_float + a2_float]
2945                + self.w_dif_float_3_line[32 - d3_float + a3_float]
2946                + self.w_dif_ground_1_line[96 + a1_ground - d1_ground]
2947                + self.w_dif_ground_2_line[64 + a2_ground - d2_ground]
2948                + self.w_dif_ground_3_line[32 + a3_ground - d3_ground])
2949            + self.w_trap_3_num[trap_3_num]
2950            + self.bias;
2951        return 1.0 / (1.0 + (-val).exp());
2952    }
2953
2954    pub fn move_average(&mut self) {
2955        let base = self.w_float_1_line[0];
2956        self.bias += base;
2957        for i in self.w_float_1_line.iter_mut() {
2958            if *i == 0.0 {
2959                continue;
2960            }
2961            *i -= base;
2962        }
2963        let base = self.w_float_2_line[0];
2964        self.bias += base;
2965        for i in self.w_float_2_line.iter_mut() {
2966            if *i == 0.0 {
2967                continue;
2968            }
2969            *i -= base;
2970        }
2971        let base = self.w_float_3_line[0];
2972        self.bias += base;
2973        for i in self.w_float_3_line.iter_mut() {
2974            if *i == 0.0 {
2975                continue;
2976            }
2977            *i -= base;
2978        }
2979        let base = self.w_dif_float_1_line[96];
2980        self.bias += base;
2981        for i in self.w_dif_float_1_line.iter_mut() {
2982            if *i == 0.0 {
2983                continue;
2984            }
2985            *i -= base;
2986        }
2987        let base = self.w_dif_float_2_line[64];
2988        self.bias += base;
2989        for i in self.w_dif_float_2_line.iter_mut() {
2990            if *i == 0.0 {
2991                continue;
2992            }
2993            *i -= base;
2994        }
2995        let base = self.w_dif_float_3_line[32];
2996        self.bias += base;
2997        for i in self.w_dif_float_3_line.iter_mut() {
2998            if *i == 0.0 {
2999                continue;
3000            }
3001            *i -= base;
3002        }
3003
3004        let base = self.w_ground_1_line[0];
3005        self.bias += base;
3006        for i in self.w_ground_1_line.iter_mut() {
3007            if *i == 0.0 {
3008                continue;
3009            }
3010            *i -= base;
3011        }
3012        let base = self.w_ground_2_line[0];
3013        self.bias += base;
3014        for i in self.w_ground_2_line.iter_mut() {
3015            if *i == 0.0 {
3016                continue;
3017            }
3018            *i -= base;
3019        }
3020        let base = self.w_ground_3_line[0];
3021        self.bias += base;
3022        for i in self.w_ground_3_line.iter_mut() {
3023            if *i == 0.0 {
3024                continue;
3025            }
3026            *i -= base;
3027        }
3028        let base = self.w_dif_ground_1_line[96];
3029        self.bias += base;
3030        for i in self.w_dif_ground_1_line.iter_mut() {
3031            if *i == 0.0 {
3032                continue;
3033            }
3034            *i -= base;
3035        }
3036        let base = self.w_dif_ground_2_line[64];
3037        self.bias += base;
3038        for i in self.w_dif_ground_2_line.iter_mut() {
3039            if *i == 0.0 {
3040                continue;
3041            }
3042            *i -= base;
3043        }
3044        let base = self.w_dif_ground_3_line[32];
3045        self.bias += base;
3046        for i in self.w_dif_ground_3_line.iter_mut() {
3047            if *i == 0.0 {
3048                continue;
3049            }
3050            *i -= base;
3051        }
3052
3053        let base = self.w_pos_1_line[0];
3054        self.bias += base;
3055        for i in self.w_pos_1_line.iter_mut() {
3056            if *i == 0.0 {
3057                continue;
3058            }
3059            *i -= base;
3060        }
3061        let base = self.w_pos_2_line[0];
3062        self.bias += base;
3063        for i in self.w_pos_2_line.iter_mut() {
3064            if *i == 0.0 {
3065                continue;
3066            }
3067            *i -= base;
3068        }
3069        let base = self.w_pos_3_line[0];
3070        self.bias += base;
3071        for i in self.w_pos_3_line.iter_mut() {
3072            if *i == 0.0 {
3073                continue;
3074            }
3075            *i -= base;
3076        }
3077        let base = self.w_dpos_1_line[0];
3078        self.bias += base;
3079        for i in self.w_dpos_1_line.iter_mut() {
3080            if *i == 0.0 {
3081                continue;
3082            }
3083            *i -= base;
3084        }
3085        let base = self.w_dpos_2_line[0];
3086        self.bias += base;
3087        for i in self.w_dpos_2_line.iter_mut() {
3088            if *i == 0.0 {
3089                continue;
3090            }
3091            *i -= base;
3092        }
3093        let base = self.w_dpos_3_line[0];
3094        self.bias += base;
3095        for i in self.w_dpos_3_line.iter_mut() {
3096            if *i == 0.0 {
3097                continue;
3098            }
3099            *i -= base;
3100        }
3101
3102        let base = self.w_bf_1_line[0];
3103        self.bias += base;
3104        for i in self.w_bf_1_line.iter_mut() {
3105            if *i == 0.0 {
3106                continue;
3107            }
3108            *i -= base;
3109        }
3110        let base = self.w_bf_2_line[0];
3111        self.bias += base;
3112        for i in self.w_bf_2_line.iter_mut() {
3113            if *i == 0.0 {
3114                continue;
3115            }
3116            *i -= base;
3117        }
3118        let base = self.w_bf_3_line[0];
3119        self.bias += base;
3120        for i in self.w_bf_3_line.iter_mut() {
3121            if *i == 0.0 {
3122                continue;
3123            }
3124            *i -= base;
3125        }
3126        let base = self.w_dbf_1_line[0];
3127        self.bias += base;
3128        for i in self.w_dbf_1_line.iter_mut() {
3129            if *i == 0.0 {
3130                continue;
3131            }
3132            *i -= base;
3133        }
3134        let base = self.w_dbf_2_line[0];
3135        self.bias += base;
3136        for i in self.w_dbf_2_line.iter_mut() {
3137            if *i == 0.0 {
3138                continue;
3139            }
3140            *i -= base;
3141        }
3142        let base = self.w_dbf_3_line[0];
3143        self.bias += base;
3144        for i in self.w_dbf_3_line.iter_mut() {
3145            if *i == 0.0 {
3146                continue;
3147            }
3148            *i -= base;
3149        }
3150
3151        let base = self.w_bg_1_line[0];
3152        self.bias += base;
3153        for i in self.w_bg_1_line.iter_mut() {
3154            if *i == 0.0 {
3155                continue;
3156            }
3157            *i -= base;
3158        }
3159        let base = self.w_bg_2_line[0];
3160        self.bias += base;
3161        for i in self.w_bg_2_line.iter_mut() {
3162            if *i == 0.0 {
3163                continue;
3164            }
3165            *i -= base;
3166        }
3167        let base = self.w_bg_3_line[0];
3168        self.bias += base;
3169        for i in self.w_bg_3_line.iter_mut() {
3170            if *i == 0.0 {
3171                continue;
3172            }
3173            *i -= base;
3174        }
3175        let base = self.w_dbg_1_line[0];
3176        self.bias += base;
3177        for i in self.w_dbg_1_line.iter_mut() {
3178            if *i == 0.0 {
3179                continue;
3180            }
3181            *i -= base;
3182        }
3183        let base = self.w_dbg_2_line[0];
3184        self.bias += base;
3185        for i in self.w_dbg_2_line.iter_mut() {
3186            if *i == 0.0 {
3187                continue;
3188            }
3189            *i -= base;
3190        }
3191        let base = self.w_dbg_3_line[0];
3192        self.bias += base;
3193        for i in self.w_dbg_3_line.iter_mut() {
3194            if *i == 0.0 {
3195                continue;
3196            }
3197            *i -= base;
3198        }
3199
3200        let base = self.w_bf_1_line[0];
3201        self.bias += base;
3202        for i in self.w_bf_1_line.iter_mut() {
3203            if *i == 0.0 {
3204                continue;
3205            }
3206            *i -= base;
3207        }
3208        let base = self.w_bf_2_line[0];
3209        self.bias += base;
3210        for i in self.w_bf_2_line.iter_mut() {
3211            if *i == 0.0 {
3212                continue;
3213            }
3214            *i -= base;
3215        }
3216        let base = self.w_bf_3_line[0];
3217        self.bias += base;
3218        for i in self.w_bf_3_line.iter_mut() {
3219            if *i == 0.0 {
3220                continue;
3221            }
3222            *i -= base;
3223        }
3224        let base = self.w_bg_1_line[0];
3225        self.bias += base;
3226        for i in self.w_bg_1_line.iter_mut() {
3227            if *i == 0.0 {
3228                continue;
3229            }
3230            *i -= base;
3231        }
3232        let base = self.w_bg_2_line[0];
3233        self.bias += base;
3234        for i in self.w_bg_2_line.iter_mut() {
3235            if *i == 0.0 {
3236                continue;
3237            }
3238            *i -= base;
3239        }
3240        let base = self.w_bg_3_line[0];
3241        self.bias += base;
3242        for i in self.w_bg_3_line.iter_mut() {
3243            if *i == 0.0 {
3244                continue;
3245            }
3246            *i -= base;
3247        }
3248
3249        let base = self.w_pos[0];
3250        self.bias += base;
3251        for i in self.w_pos.iter_mut() {
3252            if *i == 0.0 {
3253                continue;
3254            }
3255            *i -= base;
3256        }
3257        let base = self.w_dpos[0];
3258        self.bias += base;
3259        for i in self.w_dpos.iter_mut() {
3260            if *i == 0.0 {
3261                continue;
3262            }
3263            *i -= base;
3264        }
3265        let base = self.w_trap_3_num[0];
3266        self.bias += base;
3267        for i in self.w_trap_3_num.iter_mut() {
3268            if *i == 0.0 {
3269                continue;
3270            }
3271            *i -= base;
3272        }
3273    }
3274
3275    pub fn load(&mut self, file: String) -> Result<()> {
3276        use std::fs::File;
3277        use std::io::{BufRead, BufReader};
3278        let mut line = BufReader::new(File::open(file)?).lines();
3279
3280        let _ = line.next();
3281        for i in 0..96 {
3282            let l = line.next().unwrap()?;
3283            let num: f32 = l.parse().unwrap();
3284            self.w_float_1_line[i] = num;
3285        }
3286        let _ = line.next();
3287        for i in 0..96 {
3288            let l = line.next().unwrap()?;
3289            let num: f32 = l.parse().unwrap();
3290            self.w_ground_1_line[i] = num;
3291        }
3292        let _ = line.next();
3293        for i in 0..64 {
3294            let l = line.next().unwrap()?;
3295            let num: f32 = l.parse().unwrap();
3296            self.w_float_2_line[i] = num;
3297        }
3298        let _ = line.next();
3299        for i in 0..64 {
3300            let l = line.next().unwrap()?;
3301            let num: f32 = l.parse().unwrap();
3302            self.w_ground_2_line[i] = num;
3303        }
3304        let _ = line.next();
3305        for i in 0..32 {
3306            let l = line.next().unwrap()?;
3307            let num: f32 = l.parse().unwrap();
3308            self.w_float_3_line[i] = num;
3309        }
3310        let _ = line.next();
3311        for i in 0..32 {
3312            let l = line.next().unwrap()?;
3313            let num: f32 = l.parse().unwrap();
3314            self.w_ground_3_line[i] = num;
3315        }
3316
3317        let _ = line.next();
3318        for i in 0..192 {
3319            let l = line.next().unwrap()?;
3320            let num: f32 = l.parse().unwrap();
3321            self.w_dif_float_1_line[i] = num;
3322        }
3323        let _ = line.next();
3324        for i in 0..192 {
3325            let l = line.next().unwrap()?;
3326            let num: f32 = l.parse().unwrap();
3327            self.w_dif_ground_1_line[i] = num;
3328        }
3329        let _ = line.next();
3330        for i in 0..128 {
3331            let l = line.next().unwrap()?;
3332            let num: f32 = l.parse().unwrap();
3333            self.w_dif_float_2_line[i] = num;
3334        }
3335        let _ = line.next();
3336        for i in 0..128 {
3337            let l = line.next().unwrap()?;
3338            let num: f32 = l.parse().unwrap();
3339            self.w_dif_ground_2_line[i] = num;
3340        }
3341        let _ = line.next();
3342        for i in 0..64 {
3343            let l = line.next().unwrap()?;
3344            let num: f32 = l.parse().unwrap();
3345            self.w_dif_float_3_line[i] = num;
3346        }
3347        let _ = line.next();
3348        for i in 0..64 {
3349            let l = line.next().unwrap()?;
3350            let num: f32 = l.parse().unwrap();
3351            self.w_dif_ground_3_line[i] = num;
3352        }
3353
3354        let _ = line.next();
3355        for i in 0..64 {
3356            let l = line.next().unwrap()?;
3357            let num: f32 = l.parse().unwrap();
3358            self.w_pos_1_line[i] = num;
3359        }
3360        let _ = line.next();
3361        for i in 0..64 {
3362            let l = line.next().unwrap()?;
3363            let num: f32 = l.parse().unwrap();
3364            self.w_pos_2_line[i] = num;
3365        }
3366        let _ = line.next();
3367        for i in 0..64 {
3368            let l = line.next().unwrap()?;
3369            let num: f32 = l.parse().unwrap();
3370            self.w_pos_3_line[i] = num;
3371        }
3372        let _ = line.next();
3373        for i in 0..64 {
3374            let l = line.next().unwrap()?;
3375            let num: f32 = l.parse().unwrap();
3376            self.w_dpos_1_line[i] = num;
3377        }
3378        let _ = line.next();
3379        for i in 0..64 {
3380            let l = line.next().unwrap()?;
3381            let num: f32 = l.parse().unwrap();
3382            self.w_dpos_2_line[i] = num;
3383        }
3384        let _ = line.next();
3385        for i in 0..64 {
3386            let l = line.next().unwrap()?;
3387            let num: f32 = l.parse().unwrap();
3388            self.w_dpos_3_line[i] = num;
3389        }
3390
3391        let _ = line.next();
3392        for i in 0..64 {
3393            let l = line.next().unwrap()?;
3394            let num: f32 = l.parse().unwrap();
3395            self.w_bf_1_line[i] = num;
3396        }
3397        let _ = line.next();
3398        for i in 0..64 {
3399            let l = line.next().unwrap()?;
3400            let num: f32 = l.parse().unwrap();
3401            self.w_bf_2_line[i] = num;
3402        }
3403        let _ = line.next();
3404        for i in 0..64 {
3405            let l = line.next().unwrap()?;
3406            let num: f32 = l.parse().unwrap();
3407            self.w_bf_3_line[i] = num;
3408        }
3409        let _ = line.next();
3410        for i in 0..64 {
3411            let l = line.next().unwrap()?;
3412            let num: f32 = l.parse().unwrap();
3413            self.w_bg_1_line[i] = num;
3414        }
3415        let _ = line.next();
3416        for i in 0..64 {
3417            let l = line.next().unwrap()?;
3418            let num: f32 = l.parse().unwrap();
3419            self.w_bg_2_line[i] = num;
3420        }
3421        let _ = line.next();
3422        for i in 0..64 {
3423            let l = line.next().unwrap()?;
3424            let num: f32 = l.parse().unwrap();
3425            self.w_bg_3_line[i] = num;
3426        }
3427
3428        let _ = line.next();
3429        for i in 0..64 {
3430            let l = line.next().unwrap()?;
3431            let num: f32 = l.parse().unwrap();
3432            self.w_dbf_1_line[i] = num;
3433        }
3434        let _ = line.next();
3435        for i in 0..64 {
3436            let l = line.next().unwrap()?;
3437            let num: f32 = l.parse().unwrap();
3438            self.w_dbf_2_line[i] = num;
3439        }
3440        let _ = line.next();
3441        for i in 0..64 {
3442            let l = line.next().unwrap()?;
3443            let num: f32 = l.parse().unwrap();
3444            self.w_dbf_3_line[i] = num;
3445        }
3446        let _ = line.next();
3447        for i in 0..64 {
3448            let l = line.next().unwrap()?;
3449            let num: f32 = l.parse().unwrap();
3450            self.w_dbg_1_line[i] = num;
3451        }
3452        let _ = line.next();
3453        for i in 0..64 {
3454            let l = line.next().unwrap()?;
3455            let num: f32 = l.parse().unwrap();
3456            self.w_dbg_2_line[i] = num;
3457        }
3458        let _ = line.next();
3459        for i in 0..64 {
3460            let l = line.next().unwrap()?;
3461            let num: f32 = l.parse().unwrap();
3462            self.w_dbg_3_line[i] = num;
3463        }
3464
3465        let _ = line.next();
3466        for i in 0..64 {
3467            let l = line.next().unwrap()?;
3468            let num: f32 = l.parse().unwrap();
3469            self.w_pos[i] = num;
3470        }
3471        let _ = line.next();
3472        for i in 0..64 {
3473            let l = line.next().unwrap()?;
3474            let num: f32 = l.parse().unwrap();
3475            self.w_dpos[i] = num;
3476        }
3477
3478        let _ = line.next();
3479        for i in 0..13 {
3480            let l = line.next().unwrap()?;
3481            let num: f32 = l.parse().unwrap();
3482            self.w_trap_3_num[i] = num;
3483        }
3484
3485        let l = line.next().unwrap()?;
3486        let bias = l.parse().unwrap();
3487        self.bias = bias;
3488
3489        Ok(())
3490    }
3491
3492    pub fn save(&self, file: String) -> Result<()> {
3493        use std::fs::File;
3494        use std::{
3495            io,
3496            io::{BufRead, BufReader, Write},
3497        };
3498        let mut file = File::create(file)?;
3499
3500        writeln!(file, "w_float_1_line");
3501        for i in 0..96 {
3502            writeln!(file, "{}", self.w_float_1_line[i]);
3503        }
3504        writeln!(file, "w_ground_1_line");
3505        for i in 0..96 {
3506            writeln!(file, "{}", self.w_ground_1_line[i]);
3507        }
3508        writeln!(file, "w_float_2_line");
3509        for i in 0..64 {
3510            writeln!(file, "{}", self.w_float_2_line[i]);
3511        }
3512        writeln!(file, "w_ground_2_line");
3513        for i in 0..64 {
3514            writeln!(file, "{}", self.w_ground_2_line[i]);
3515        }
3516        writeln!(file, "w_float_3_line");
3517        for i in 0..32 {
3518            writeln!(file, "{}", self.w_float_3_line[i]);
3519        }
3520        writeln!(file, "w_ground_3_line");
3521        for i in 0..32 {
3522            writeln!(file, "{}", self.w_ground_3_line[i]);
3523        }
3524
3525        writeln!(file, "w_dif_float_1_line");
3526        for i in 0..192 {
3527            writeln!(file, "{}", self.w_dif_float_1_line[i]);
3528        }
3529        writeln!(file, "w_dif_ground_1_line");
3530        for i in 0..192 {
3531            writeln!(file, "{}", self.w_dif_ground_1_line[i]);
3532        }
3533        writeln!(file, "w_dif_float_2_line");
3534        for i in 0..128 {
3535            writeln!(file, "{}", self.w_dif_float_2_line[i]);
3536        }
3537        writeln!(file, "w_dif_ground_2_line");
3538        for i in 0..128 {
3539            writeln!(file, "{}", self.w_dif_ground_2_line[i]);
3540        }
3541        writeln!(file, "w_dif_float_3_line");
3542        for i in 0..64 {
3543            writeln!(file, "{}", self.w_dif_float_3_line[i]);
3544        }
3545        writeln!(file, "w_dif_ground_3_line");
3546        for i in 0..64 {
3547            writeln!(file, "{}", self.w_dif_ground_3_line[i]);
3548        }
3549
3550        writeln!(file, "w_pos_1_line");
3551        for i in 0..64 {
3552            writeln!(file, "{}", self.w_pos_1_line[i]);
3553        }
3554        writeln!(file, "w_pos_2_line");
3555        for i in 0..64 {
3556            writeln!(file, "{}", self.w_pos_2_line[i]);
3557        }
3558        writeln!(file, "w_pos_3_line");
3559        for i in 0..64 {
3560            writeln!(file, "{}", self.w_pos_3_line[i]);
3561        }
3562        writeln!(file, "w_dpos_1_line");
3563        for i in 0..64 {
3564            writeln!(file, "{}", self.w_dpos_1_line[i]);
3565        }
3566        writeln!(file, "w_dpos_2_line");
3567        for i in 0..64 {
3568            writeln!(file, "{}", self.w_dpos_2_line[i]);
3569        }
3570        writeln!(file, "w_dpos_3_line");
3571        for i in 0..64 {
3572            writeln!(file, "{}", self.w_dpos_3_line[i]);
3573        }
3574
3575        writeln!(file, "w_bf_1_line");
3576        for i in 0..64 {
3577            writeln!(file, "{}", self.w_bf_1_line[i]);
3578        }
3579        writeln!(file, "w_bf_2_line");
3580        for i in 0..64 {
3581            writeln!(file, "{}", self.w_bf_2_line[i]);
3582        }
3583        writeln!(file, "w_bf_3_line");
3584        for i in 0..64 {
3585            writeln!(file, "{}", self.w_bf_3_line[i]);
3586        }
3587        writeln!(file, "w_bg_1_line");
3588        for i in 0..64 {
3589            writeln!(file, "{}", self.w_bg_1_line[i]);
3590        }
3591        writeln!(file, "w_bg_2_line");
3592        for i in 0..64 {
3593            writeln!(file, "{}", self.w_bg_2_line[i]);
3594        }
3595        writeln!(file, "w_bg_3_line");
3596        for i in 0..64 {
3597            writeln!(file, "{}", self.w_bg_3_line[i]);
3598        }
3599
3600        writeln!(file, "w_dbf_1_line");
3601        for i in 0..64 {
3602            writeln!(file, "{}", self.w_dbf_1_line[i]);
3603        }
3604        writeln!(file, "w_dbf_2_line");
3605        for i in 0..64 {
3606            writeln!(file, "{}", self.w_dbf_2_line[i]);
3607        }
3608        writeln!(file, "w_dbf_3_line");
3609        for i in 0..64 {
3610            writeln!(file, "{}", self.w_dbf_3_line[i]);
3611        }
3612        writeln!(file, "w_dbg_1_line");
3613        for i in 0..64 {
3614            writeln!(file, "{}", self.w_dbg_1_line[i]);
3615        }
3616        writeln!(file, "w_dbg_2_line");
3617        for i in 0..64 {
3618            writeln!(file, "{}", self.w_dbg_2_line[i]);
3619        }
3620        writeln!(file, "w_dbg_3_line");
3621        for i in 0..64 {
3622            writeln!(file, "{}", self.w_dbg_3_line[i]);
3623        }
3624
3625        writeln!(file, "w_pos");
3626        for i in 0..64 {
3627            writeln!(file, "{}", self.w_pos[i]);
3628        }
3629        writeln!(file, "w_dpos");
3630        for i in 0..64 {
3631            writeln!(file, "{}", self.w_dpos[i]);
3632        }
3633
3634        writeln!(file, "w_trap_3_num");
3635        for i in 0..13 {
3636            writeln!(file, "{}", self.w_trap_3_num[i]);
3637        }
3638
3639        writeln!(file, "{}", self.bias);
3640
3641        file.flush()?;
3642        Ok(())
3643    }
3644}
3645
3646impl EvaluatorF for LineEvaluator {
3647    fn eval_func_f32(&self, b: &Board) -> f32 {
3648        return self.evaluate_board(b);
3649    }
3650}
3651
3652pub trait Trainable {
3653    fn update(&mut self, b: &Board, delta: f32) {}
3654    fn get_val(&self, b: &Board) -> f32;
3655    fn save(&self, file: String) -> Result<()>;
3656    fn load(&mut self, file: String) -> Result<()>;
3657    fn eval(&mut self) {}
3658    fn train(&mut self) {}
3659}
3660
3661#[derive(Clone)]
3662pub struct LELearningParam {
3663    pub wfl1: bool,
3664    pub wfl2: bool,
3665    pub wfl3: bool,
3666    pub wgl1: bool,
3667    pub wgl2: bool,
3668    pub wgl3: bool,
3669    pub wdfl1: bool,
3670    pub wdfl2: bool,
3671    pub wdfl3: bool,
3672    pub wdgl1: bool,
3673    pub wdgl2: bool,
3674    pub wdgl3: bool,
3675    pub wpl1: bool,
3676    pub wpl2: bool,
3677    pub wpl3: bool,
3678    pub wdpl1: bool,
3679    pub wdpl2: bool,
3680    pub wdpl3: bool,
3681    pub wbfl1: bool,
3682    pub wbfl2: bool,
3683    pub wbfl3: bool,
3684    pub wbgl1: bool,
3685    pub wbgl2: bool,
3686    pub wbgl3: bool,
3687    pub wdbfl1: bool,
3688    pub wdbfl2: bool,
3689    pub wdbfl3: bool,
3690    pub wdbgl1: bool,
3691    pub wdbgl2: bool,
3692    pub wdbgl3: bool,
3693    pub wp: bool,
3694    pub wdp: bool,
3695    pub wtn3: bool,
3696    pub bias: bool,
3697}
3698
3699impl LELearningParam {
3700    pub fn new() -> Self {
3701        LELearningParam {
3702            wfl1: true,
3703            wfl2: true,
3704            wfl3: true,
3705            wgl1: true,
3706            wgl2: true,
3707            wgl3: true,
3708            wdfl1: true,
3709            wdfl2: true,
3710            wdfl3: true,
3711            wdgl1: true,
3712            wdgl2: true,
3713            wdgl3: true,
3714            wpl1: true,
3715            wpl2: true,
3716            wpl3: true,
3717            wdpl1: true,
3718            wdpl2: true,
3719            wdpl3: true,
3720            wbfl1: true,
3721            wbfl2: true,
3722            wbfl3: true,
3723            wbgl1: true,
3724            wbgl2: true,
3725            wbgl3: true,
3726            wdbfl1: true,
3727            wdbfl2: true,
3728            wdbfl3: true,
3729            wdbgl1: true,
3730            wdbgl2: true,
3731            wdbgl3: true,
3732            wp: true,
3733            wdp: true,
3734            wtn3: true,
3735            bias: true,
3736        }
3737    }
3738    pub fn from(hash: u64) -> Self {
3739        LELearningParam {
3740            wfl1: (hash >> 0) & 1 == 1,
3741            wfl2: (hash >> 1) & 1 == 1,
3742            wfl3: (hash >> 2) & 1 == 1,
3743            wgl1: (hash >> 3) & 1 == 1,
3744            wgl2: (hash >> 4) & 1 == 1,
3745            wgl3: (hash >> 5) & 1 == 1,
3746
3747            wdfl1: (hash >> 6) & 1 == 1,
3748            wdfl2: (hash >> 7) & 1 == 1,
3749            wdfl3: (hash >> 8) & 1 == 1,
3750            wdgl1: (hash >> 9) & 1 == 1,
3751            wdgl2: (hash >> 10) & 1 == 1,
3752            wdgl3: (hash >> 11) & 1 == 1,
3753
3754            wpl1: (hash >> 12) & 1 == 1,
3755            wpl2: (hash >> 13) & 1 == 1,
3756            wpl3: (hash >> 14) & 1 == 1,
3757            wdpl1: (hash >> 15) & 1 == 1,
3758            wdpl2: (hash >> 16) & 1 == 1,
3759            wdpl3: (hash >> 17) & 1 == 1,
3760
3761            wbfl1: (hash >> 18) & 1 == 1,
3762            wbfl2: (hash >> 19) & 1 == 1,
3763            wbfl3: (hash >> 20) & 1 == 1,
3764            wbgl1: (hash >> 21) & 1 == 1,
3765            wbgl2: (hash >> 22) & 1 == 1,
3766            wbgl3: (hash >> 23) & 1 == 1,
3767
3768            wdbfl1: (hash >> 24) & 1 == 1,
3769            wdbfl2: (hash >> 25) & 1 == 1,
3770            wdbfl3: (hash >> 26) & 1 == 1,
3771            wdbgl1: (hash >> 27) & 1 == 1,
3772            wdbgl2: (hash >> 28) & 1 == 1,
3773            wdbgl3: (hash >> 29) & 1 == 1,
3774
3775            wp: (hash >> 30) & 1 == 1,
3776            wdp: (hash >> 31) & 1 == 1,
3777
3778            wtn3: (hash >> 32) & 1 == 1,
3779            bias: (hash >> 33) & 1 == 1,
3780        }
3781    }
3782}
3783
3784#[derive(Clone)]
3785pub struct TrainableLineEvaluator {
3786    main: LineEvaluator,
3787    v: LineEvaluator,
3788    m: LineEvaluator,
3789    lr: f32,
3790    pub param: LELearningParam,
3791}
3792
3793impl TrainableLineEvaluator {
3794    pub fn new(lr: f32) -> Self {
3795        TrainableLineEvaluator {
3796            main: LineEvaluator::new(),
3797            v: LineEvaluator::new(),
3798            m: LineEvaluator::new(),
3799            lr: lr,
3800            param: LELearningParam::new(),
3801        }
3802    }
3803
3804    pub fn from(e: LineEvaluator, lr: f32) -> Self {
3805        TrainableLineEvaluator {
3806            main: e,
3807            v: LineEvaluator::new(),
3808            m: LineEvaluator::new(),
3809            lr: lr,
3810            param: LELearningParam::new(),
3811        }
3812    }
3813
3814    pub fn set_param(&mut self, hash: u64) {
3815        self.param = LELearningParam::from(hash);
3816    }
3817}
3818
3819impl Trainable for TrainableLineEvaluator {
3820    fn update(&mut self, b: &Board, delta: f32) {
3821        let (
3822            a1,
3823            a2,
3824            a3,
3825            a1_,
3826            a2_,
3827            a3_,
3828            d1,
3829            d2,
3830            d3,
3831            d1_,
3832            d2_,
3833            d3_,
3834            af1_mask,
3835            af2_mask,
3836            af3_mask,
3837            ag1_mask,
3838            ag2_mask,
3839            ag3_mask,
3840            df1_mask,
3841            df2_mask,
3842            df3_mask,
3843            dg1_mask,
3844            dg2_mask,
3845            dg3_mask,
3846            pa1_mask,
3847            pa2_mask,
3848            pa3_mask,
3849            pd1_mask,
3850            pd2_mask,
3851            pd3_mask,
3852            trap_3_num,
3853            att,
3854            def,
3855        ) = self.main.get_counts(b);
3856        // とりあえずsgd
3857        let val = self.main.evaluate_board(b);
3858        let dv = val * (1.0 - val);
3859        let delta = self.lr * delta * dv;
3860        if self.param.wfl1 {
3861            self.main.w_float_1_line[a1] += delta;
3862        }
3863        if self.param.wfl2 {
3864            self.main.w_float_2_line[a2] += delta;
3865        }
3866        if self.param.wfl3 {
3867            self.main.w_float_3_line[a3] += delta;
3868        }
3869        if self.param.wgl1 {
3870            self.main.w_ground_1_line[a1_] += delta;
3871        }
3872        if self.param.wgl2 {
3873            self.main.w_ground_2_line[a2_] += delta;
3874        }
3875        if self.param.wgl3 {
3876            self.main.w_ground_3_line[a3_] += delta;
3877        }
3878        if self.param.wdfl1 {
3879            self.main.w_dif_float_1_line[96 + a1 - d1] += delta;
3880        }
3881        if self.param.wdfl2 {
3882            self.main.w_dif_float_2_line[64 + a2 - d2] += delta;
3883        }
3884        if self.param.wdfl3 {
3885            self.main.w_dif_float_3_line[32 + a3 - d3] += delta;
3886        }
3887        if self.param.wdgl1 {
3888            self.main.w_dif_ground_1_line[96 + a1_ - d1_] += delta;
3889        }
3890        if self.param.wdgl2 {
3891            self.main.w_dif_ground_2_line[64 + a2_ - d2_] += delta;
3892        }
3893        if self.param.wdgl3 {
3894            self.main.w_dif_ground_3_line[32 + a3_ - d3_] += delta;
3895        }
3896        if self.param.wtn3 {
3897            self.main.w_trap_3_num[trap_3_num] += delta;
3898        }
3899        if self.param.bias {
3900            self.main.bias += delta;
3901        }
3902
3903        for i in 0..64 {
3904            let bit = 1 << i;
3905            if af1_mask & bit != 0 && self.param.wbfl1 {
3906                self.main.w_bf_1_line[i] += delta;
3907            }
3908            if af2_mask & bit != 0 && self.param.wbfl2 {
3909                self.main.w_bf_2_line[i] += delta;
3910            }
3911            if af3_mask & bit != 0 && self.param.wbfl3 {
3912                self.main.w_bf_3_line[i] += delta;
3913            }
3914            if df1_mask & bit != 0 && self.param.wdbfl1 {
3915                self.main.w_dbf_1_line[i] += delta;
3916            }
3917            if df2_mask & bit != 0 && self.param.wdbfl2 {
3918                self.main.w_dbf_2_line[i] += delta;
3919            }
3920            if df3_mask & bit != 0 && self.param.wdbfl3 {
3921                self.main.w_dbf_3_line[i] += delta;
3922            }
3923            if ag1_mask & bit != 0 && self.param.wbgl1 {
3924                self.main.w_bg_1_line[i] += delta;
3925            }
3926            if ag2_mask & bit != 0 && self.param.wbgl2 {
3927                self.main.w_bg_2_line[i] += delta;
3928            }
3929            if ag3_mask & bit != 0 && self.param.wbgl3 {
3930                self.main.w_bg_3_line[i] += delta;
3931            }
3932            if dg1_mask & bit != 0 && self.param.wdbgl1 {
3933                self.main.w_dbg_1_line[i] += delta;
3934            }
3935            if dg2_mask & bit != 0 && self.param.wdbgl2 {
3936                self.main.w_dbg_2_line[i] += delta;
3937            }
3938            if dg3_mask & bit != 0 && self.param.wdbgl3 {
3939                self.main.w_dbg_3_line[i] += delta;
3940            }
3941            if pa1_mask & bit != 0 && self.param.wpl1 {
3942                self.main.w_pos_1_line[i] += delta;
3943            } else if pd1_mask & bit != 0 && self.param.wdpl1 {
3944                self.main.w_dpos_1_line[i] += delta;
3945            }
3946            if pa2_mask & bit != 0 && self.param.wpl2 {
3947                self.main.w_pos_2_line[i] += delta;
3948            } else if pd2_mask & bit != 0 && self.param.wdpl2 {
3949                self.main.w_dpos_2_line[i] += delta;
3950            }
3951            if pa3_mask & bit != 0 && self.param.wpl3 {
3952                self.main.w_pos_3_line[i] += delta;
3953            } else if pd3_mask & bit != 0 && self.param.wdpl3 {
3954                self.main.w_dpos_3_line[i] += delta;
3955            }
3956            if att & bit != 0 && self.param.wp {
3957                self.main.w_pos[i] += delta;
3958            } else if def & bit != 0 && self.param.wdp {
3959                self.main.w_dpos[i] += delta;
3960            }
3961        }
3962    }
3963
3964    fn get_val(&self, b: &Board) -> f32 {
3965        self.main.evaluate_board(b)
3966    }
3967
3968    fn save(&self, file: String) -> Result<()> {
3969        self.main.save(file)
3970    }
3971    fn load(&mut self, file: String) -> Result<()> {
3972        self.main.load(file)
3973    }
3974    fn eval(&mut self) {}
3975    fn train(&mut self) {
3976        self.main.move_average();
3977        return;
3978        print!("wf1:[");
3979        for f in self.main.w_float_1_line.iter() {
3980            print!("{},", f);
3981        }
3982        println!("]");
3983        print!("wf2:[");
3984        for f in self.main.w_float_2_line.iter() {
3985            print!("{},", f);
3986        }
3987        println!("]");
3988        print!("wf3:[");
3989        for f in self.main.w_float_3_line.iter() {
3990            print!("{},", f);
3991        }
3992        println!("]");
3993        print!("wg1:[");
3994        for f in self.main.w_ground_1_line.iter() {
3995            print!("{},", f);
3996        }
3997        println!("]");
3998        print!("wg2:[");
3999        for f in self.main.w_ground_2_line.iter() {
4000            print!("{},", f);
4001        }
4002        println!("]");
4003        print!("wg3:[");
4004        for f in self.main.w_ground_3_line.iter() {
4005            print!("{},", f);
4006        }
4007        println!("]");
4008
4009        println!("bias:{}", self.main.bias);
4010    }
4011}
4012
4013impl EvaluatorF for TrainableLineEvaluator {
4014    fn eval_func_f32(&self, b: &Board) -> f32 {
4015        return self.main.evaluate_board(b).clamp(0.0, 1.0);
4016    }
4017}
4018
4019pub struct RandomEvaluator {}
4020
4021impl RandomEvaluator {
4022    pub fn new() -> Self {
4023        return RandomEvaluator {};
4024    }
4025}
4026
4027impl Evaluator for RandomEvaluator {
4028    fn eval_func(&self, b: &Board) -> i32 {
4029        return (get_random_usize() % 1300) as i32;
4030    }
4031}
4032
4033pub struct MateWrapperActor {
4034    main_actor: Box<dyn GetAction>,
4035}
4036
4037impl MateWrapperActor {
4038    pub fn new(actor: Box<dyn GetAction>) -> Self {
4039        return MateWrapperActor { main_actor: actor };
4040    }
4041}
4042
4043impl GetAction for MateWrapperActor {
4044    fn get_action(&self, b: &Board) -> u8 {
4045        use proconio::input;
4046        let res = mate_check_horizontal(b);
4047        if let Some((flag, action)) = res {
4048            if cfg!(any(feature = "render", feature = "view")) {
4049                if flag {
4050                    println!("mate -> {action}");
4051                } else {
4052                    println!("not mate -> {action}");
4053                }
4054            }
4055            // println!("->[{action}]");
4056            return action;
4057        } else {
4058            let action = self.main_actor.get_action(b);
4059            return action;
4060        }
4061    }
4062}
4063
4064pub struct MateNegAlpha {
4065    main_eval: Box<dyn Evaluator>,
4066    depth: u8,
4067}
4068
4069impl MateNegAlpha {
4070    pub fn new(main_eval: Box<dyn Evaluator>, depth: u8) -> Self {
4071        return MateNegAlpha {
4072            main_eval: main_eval,
4073            depth: depth,
4074        };
4075    }
4076
4077    pub fn eval_with_negalpha(&self, b: &Board) -> (u8, i32, i32) {
4078        let res = mate_check_horizontal(b);
4079        if let Some((flag, action)) = res {
4080            if flag {
4081                return (action, MAX, 0);
4082            }
4083        }
4084        let (action, v, count) = negalpha(b, self.depth, -MAX - 1, MAX + 1, &self.main_eval);
4085
4086        return (action, v, count);
4087    }
4088}
4089impl GetAction for MateNegAlpha {
4090    fn get_action(&self, b: &Board) -> u8 {
4091        use proconio::input;
4092        let res = mate_check_horizontal(b);
4093        if let Some((flag, action)) = res {
4094            if flag {
4095                // println!("mate")
4096            }
4097            // println!("->[{action}]");
4098            return action;
4099        } else {
4100            let (action, v, count) = negalpha(b, self.depth, -MAX - 1, MAX + 1, &self.main_eval);
4101            return action;
4102        }
4103        // input! {
4104        //     action: u8
4105        // }
4106        // return action;
4107    }
4108}
4109
4110pub enum PlayoutLevel {
4111    Zero,
4112    Attack4,
4113    Defence4,
4114    MateCheck,
4115    Actor(Box<dyn GetAction>),
4116}
4117
4118pub struct PlayoutEvaluator {
4119    level: PlayoutLevel,
4120}
4121
4122impl PlayoutEvaluator {
4123    pub fn new(level: PlayoutLevel) -> Self {
4124        return PlayoutEvaluator { level: level };
4125    }
4126}
4127
4128impl EvaluatorF for PlayoutEvaluator {
4129    fn eval_func_f32(&self, b: &Board) -> f32 {
4130        let mut result = 1.0;
4131        let mut b = b.clone();
4132        loop {
4133            let action;
4134            match &self.level {
4135                PlayoutLevel::Zero => {
4136                    action = get_random(&b);
4137                }
4138                PlayoutLevel::Attack4 => {
4139                    let (att, def) = b.get_att_def();
4140                    let mask = get_reach_mask(att, def);
4141                    if mask > 0 {
4142                        action = (mask.trailing_zeros() % 16) as u8;
4143                    } else {
4144                        action = get_random(&b);
4145                    }
4146                }
4147                PlayoutLevel::Defence4 => {
4148                    let (att, def) = b.get_att_def();
4149                    let mask = get_reach_mask(att, def);
4150                    if mask > 0 {
4151                        action = (mask.trailing_zeros() % 16) as u8;
4152                    } else {
4153                        let mask = get_reach_mask(def, att);
4154                        if mask > 0 {
4155                            action = (mask.trailing_zeros() % 16) as u8;
4156                        } else {
4157                            action = get_random(&b);
4158                        }
4159                    }
4160                }
4161                PlayoutLevel::MateCheck => {
4162                    let mate = mate_check_horizontal(&b);
4163                    if let Some((_, act)) = mate {
4164                        action = act;
4165                    } else {
4166                        action = get_random(&b);
4167                    }
4168                }
4169                PlayoutLevel::Actor(actor) => {
4170                    action = actor.as_ref().get_action(&b);
4171                }
4172            }
4173            b = b.next(action);
4174
4175            if b.is_win() {
4176                return result;
4177            } else if b.is_draw() {
4178                return 0.5;
4179            }
4180            result = 1.0 - result;
4181        }
4182    }
4183}