Spark Scala中为DataFrame各分组随机抽取不同数量行的实现
Group-wise Non-Replacement Sampling with Specified Sample Sizes per Group
Got it, let's tackle this problem step by step. You need to perform non-replacement random sampling on your idDF DataFrame, where each group (identified by the ID column) has a specific number of rows to extract—defined in the planDF's datesToUse column. Here's how to implement this cleanly:
First, Let's Set Up Example Data
To make this concrete, let's create sample versions of your DataFrames:
import pandas as pd import numpy as np # Your target DataFrame to sample from idDF = pd.DataFrame({ 'ID': ['A', 'A', 'A', 'B', 'B', 'C', 'C', 'C', 'C'], 'value': np.random.randn(9) # Random values for demonstration }) # Your sampling plan DataFrame planDF = pd.DataFrame({ 'ID': ['A', 'B', 'C'], 'datesToUse': [2, 1, 3], # Number of rows to sample per ID 'totalDates': [3, 2, 4] # Total rows per ID (optional for validation) })
Method 1: Merge Sampling Plan with Target DataFrame
This approach attaches the required sample size to each row first, then samples per group:
# Merge only the necessary columns from planDF into idDF merged_df = idDF.merge(planDF[['ID', 'datesToUse']], on='ID', how='left') # Define a function to sample each group def sample_group(group): # Get the required sample size for this group required_size = group['datesToUse'].iloc[0] # Handle edge case: if required size > group size, sample all rows actual_size = min(required_size, len(group)) # Perform non-replacement sampling return group.sample(n=actual_size, replace=False) # Apply sampling to each group, and drop the temporary datesToUse column sampled_df = merged_df.groupby('ID', group_keys=False).apply(sample_group) sampled_df = sampled_df.drop('datesToUse', axis=1)
Method 2: Use a Sample Size Dictionary (More Efficient for Large Data)
If your idDF is large, merging might add overhead. Instead, map sample sizes using a dictionary:
# Create a dictionary mapping ID to its required sample size sample_size_map = planDF.set_index('ID')['datesToUse'].to_dict() def sample_group_v2(group): # Get the group's ID and its required sample size group_id = group.name required_size = sample_size_map.get(group_id, 0) # Default to 0 if ID not in plan actual_size = min(required_size, len(group)) return group.sample(n=actual_size, replace=False) # Apply sampling directly to the original idDF sampled_df_v2 = idDF.groupby('ID', group_keys=False).apply(sample_group_v2)
Key Notes & Edge Cases
- Mismatched IDs: If an ID exists in
idDFbut not inplanDF, both methods default to sampling 0 rows. Adjust thegetcall in Method 2 if you want to handle this differently (e.g., throw an error or sample all rows). - Invalid Sample Sizes: If
datesToUseis negative or larger than the group's total rows, themin()function ensures we don't throw an error—we'll just sample all available rows in the group. - Non-Replacement Guarantee: The
replace=Falseparameter insample()ensures we don't sample the same row multiple times, which aligns with your requirement.
内容的提问来源于stack exchange,提问作者fractalnature
相关产品推荐
相关产品推荐

