| 577 | yield request |
| 578 | |
| 579 | def split(self): |
| 580 | response_stream = self.stub.TrainingTestSplit(self.request_generator()) |
| 581 | # verify initialization response # TODO: don't think we need this |
| 582 | init_response = next(response_stream) |
| 583 | if not init_response.initialized: |
| 584 | raise ValueError("Failed to initialize training test split") |
| 585 | |
| 586 | self.train_iter = TrainingSetSplitIterator( |
| 587 | req_queue=self.req_queue, |
| 588 | resp_queue=self.resp_queue, |
| 589 | resp_stream=response_stream, |
| 590 | request_type=serving_pb2.RequestType.TRAINING, |
| 591 | name=self.name, |
| 592 | version=self.version, |
| 593 | model=self.model, |
| 594 | test_size=self.test_size, |
| 595 | train_size=self.train_size, |
| 596 | shuffle=self.shuffle, |
| 597 | random_state=self.random_state, |
| 598 | batch_size=self.batch_size, |
| 599 | ) |
| 600 | self.test_iter = TrainingSetSplitIterator( |
| 601 | req_queue=self.req_queue, |
| 602 | resp_queue=self.resp_queue, |
| 603 | resp_stream=response_stream, |
| 604 | request_type=serving_pb2.RequestType.TEST, |
| 605 | name=self.name, |
| 606 | version=self.version, |
| 607 | model=self.model, |
| 608 | test_size=self.test_size, |
| 609 | train_size=self.train_size, |
| 610 | shuffle=self.shuffle, |
| 611 | random_state=self.random_state, |
| 612 | batch_size=self.batch_size, |
| 613 | ) |
| 614 | |
| 615 | return self.train_iter, self.test_iter |
| 616 | |
| 617 | |
| 618 | class Dataset: |