View Javadoc
1   package net.bmahe.genetics4j.gp.program;
2   
3   import java.util.Objects;
4   import java.util.random.RandomGenerator;
5   
6   import org.apache.commons.lang3.Validate;
7   
8   import net.bmahe.genetics4j.core.chromosomes.TreeNode;
9   import net.bmahe.genetics4j.gp.Operation;
10  import net.bmahe.genetics4j.gp.OperationFactory;
11  
12  public class StdProgramGenerator implements ProgramGenerator {
13  
14  	private final ProgramHelper programHelper;
15  	private final RandomGenerator randomGenerator;
16  
17  	@SuppressWarnings({ "unchecked", "rawtypes" })
18  	private <T, U> TreeNode<Operation<T>> generate(final Program program, Class<U> acceptedType, int maxDepth,
19  			int depth) {
20  
21  		OperationFactory currentNode = depth < maxDepth - 1 && randomGenerator.nextDouble() < 0.5
22  				? programHelper.pickRandomFunction(program, acceptedType)
23  				: programHelper.pickRandomTerminal(program, acceptedType);
24  
25  		final Operation<T> currentOperation = currentNode.build(program.inputSpec());
26  		final TreeNode<Operation<T>> currentTreeNode = new TreeNode<>(currentOperation);
27  
28  		final Class[] acceptedTypes = currentNode.acceptedTypes();
29  
30  		for (int i = 0; i < acceptedTypes.length; i++) {
31  			final Class childAcceptedType = acceptedTypes[i];
32  			final TreeNode<Operation<T>> operation = generate(program, childAcceptedType, maxDepth, depth + 1);
33  
34  			currentTreeNode.addChild(operation);
35  		}
36  
37  		return currentTreeNode;
38  	}
39  
40  	public StdProgramGenerator(final ProgramHelper _programHelper, final RandomGenerator _randomGenerator) {
41  		Objects.requireNonNull(_programHelper);
42  		Objects.requireNonNull(_randomGenerator);
43  
44  		this.programHelper = _programHelper;
45  		this.randomGenerator = _randomGenerator;
46  	}
47  
48  	@SuppressWarnings("rawtypes")
49  	@Override
50  	public TreeNode<Operation<?>> generate(final Program program, final int maxDepth) {
51  		Objects.requireNonNull(program);
52  		Validate.isTrue(maxDepth > 0);
53  
54  		final OperationFactory currentNode = randomGenerator.nextDouble() < 0.98
55  				? programHelper.pickRandomFunction(program)
56  				: programHelper.pickRandomTerminal(program);
57  
58  		final Operation currentOperation = currentNode.build(program.inputSpec());
59  		final TreeNode<Operation<?>> currentTreeNode = new TreeNode<>(currentOperation);
60  
61  		Class[] acceptedTypes = currentNode.acceptedTypes();
62  
63  		for (int i = 0; i < acceptedTypes.length; i++) {
64  			final Class acceptedType = acceptedTypes[i];
65  			final TreeNode<Operation<?>> operation = generate(program, acceptedType, maxDepth, 1);
66  
67  			currentTreeNode.addChild(operation);
68  		}
69  
70  		return currentTreeNode;
71  	}
72  
73  	@Override
74  	public TreeNode<Operation<?>> generate(final Program program) {
75  		return generate(program, program.maxDepth());
76  	}
77  
78  	@Override
79  	public <T, U> TreeNode<Operation<T>> generate(final Program program, final int maxDepth, final Class<U> rootType) {
80  		Objects.requireNonNull(program);
81  		Objects.requireNonNull(rootType);
82  		Validate.isTrue(maxDepth > 0);
83  
84  		return generate(program, rootType, maxDepth, 0);
85  	}
86  }