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;
22use 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 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 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 (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 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 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 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 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 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 if depth <= 1 {
378 for action in actions.iter() {
379 let next_board = &b.next(*action);
380 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 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 = 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 return (action, Ex(1.0), count);
446 } else if next_board.is_draw() {
447 hashmap.insert(hash, Ex(0.5));
448 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 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 if let Some(fail_val) = hit {
467 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 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 = 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 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 if depth <= 1 {
584 for action in actions.iter() {
585 let next_board = &b.next(*action);
586 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 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 = 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 return (action, Ex(1.0), count);
656 } else if next_board.is_draw() {
657 hashmap.insert(hash, (Ex(0.5), gen));
658 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 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 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 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 = 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 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 = 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 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 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 let new_beta = x.min(beta);
1002 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 = x;
1024 }
1025 }
1026 } else {
1027 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 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 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 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 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 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 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 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 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 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 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 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 (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 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 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 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 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 }
4097 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 }
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}