fix(model): group DailyBatchSampler batches by datetime positions - #2323
Open
张弘 (shuxue6662-a11y) wants to merge 1 commit into
Open
fix(model): group DailyBatchSampler batches by datetime positions#2323张弘 (shuxue6662-a11y) wants to merge 1 commit into
张弘 (shuxue6662-a11y) wants to merge 1 commit into
Conversation
Author
|
@microsoft-github-policy-service agree |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Description
DailyBatchSamplerassumed that the rows of a trading day are contiguous in the data source: it computed per-day counts viagroupby("datetime").size(), converted them to start offsets withcumsum, and sliced contiguous rangesnp.arange(idx, idx + count).TSDataSampler.get_index()swaps the level labels back to(datetime, instrument)but does NOT reorder the rows, which remain instrument-major (see "Indices design" inTSDataSampler). As a result each "daily" batch actually contained one instrument across many consecutive days instead of one day's cross-section.This PR changes
DailyBatchSamplerto group the actual row positions by thedatetimelevel, so every batch contains exactly one trading day's cross-section. It also:GATs.predict()to align predictions with the sampler's iteration order (dl_test.get_index()[np.concatenate(sampler_test.batches)]), since the previousdl_test.get_index()alignment assumed the old row order;__len__to return the number of batches instead of the number of samples;TSDataSampler.get_index()docstring about the instrument-major row order;Motivation and Context
Fixes #2319
How Has This Been Tested?
python -m pytest tests/test_daily_batch_sampler.py-> 2 passedblack --check -l 120on the changed filesflake8on the changed filespytest qlib/tests/test_all_pipeline.py) -- requires downloading test data; covered by CITypes of changes